diff --git a/datasette/_pytest_plugin.py b/datasette/_pytest_plugin.py index 103c616d..587380ed 100644 --- a/datasette/_pytest_plugin.py +++ b/datasette/_pytest_plugin.py @@ -89,7 +89,8 @@ def pytest_runtest_protocol(item, nextitem): continue try: ds.close() - except Exception as e: + except Exception as e: # noqa: BLE001 + # Surfaced as a pytest warning; teardown must not fail the run item.warn( pytest.PytestUnraisableExceptionWarning( f"Error closing Datasette instance: {e!r}" diff --git a/datasette/actor_auth_cookie.py b/datasette/actor_auth_cookie.py index 368213af..7503f1d5 100644 --- a/datasette/actor_auth_cookie.py +++ b/datasette/actor_auth_cookie.py @@ -1,8 +1,10 @@ -from datasette import hookimpl -from itsdangerous import BadSignature -from datasette.utils import baseconv import time +from itsdangerous import BadSignature + +from datasette import hookimpl +from datasette.utils import baseconv + @hookimpl def actor_from_request(datasette, request): diff --git a/datasette/app.py b/datasette/app.py index 4ba5d20f..c82ea075 100644 --- a/datasette/app.py +++ b/datasette/app.py @@ -2,7 +2,8 @@ from __future__ import annotations import asyncio import contextvars -from typing import TYPE_CHECKING, Any, Dict, Iterable, List, Sequence +from collections.abc import Iterable, Sequence +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from datasette.permissions import Resource @@ -12,11 +13,10 @@ import dataclasses import datetime import functools import glob -import httpx import importlib.metadata import inspect -from itsdangerous import BadSignature import json +import logging import os import re import secrets @@ -28,90 +28,41 @@ import urllib.parse from concurrent import futures from pathlib import Path -from markupsafe import Markup, escape -from itsdangerous import URLSafeSerializer +import httpx +from itsdangerous import BadSignature, URLSafeSerializer from jinja2 import ( ChoiceLoader, Environment, FileSystemLoader, - pass_context, PrefixLoader, + pass_context, ) from jinja2.environment import Template from jinja2.exceptions import TemplateNotFound +from markupsafe import Markup, escape -from .events import Event -from .column_types import SQLiteType from . import stored_queries, write_sql -from .views import Context -from .views.database import ( - database_download, - DatabaseView, - QueryView, -) -from .views.table_create_alter import ( - DatabaseForeignKeyTargetsView, - TableAlterView, - TableCreateView, - TableForeignKeySuggestionsView, -) -from .views.execute_write import ExecuteWriteAnalyzeView, ExecuteWriteView -from .views.stored_queries import ( - QueryCreateAnalyzeView, - QueryDeleteView, - QueryDefinitionView, - QueryEditView, - GlobalQueryListView, - QueryListView, - QueryParametersView, - QueryStoreView, - QueryUpdateView, -) -from .views.index import IndexView -from .views.special import ( - JsonDataView, - PatternPortfolioView, - AutocompleteDebugView, - AuthTokenView, - ApiExplorerView, - CreateTokenView, - LogoutView, - AllowDebugView, - PermissionsDebugView, - MessagesDebugView, - AllowedResourcesView, - PermissionRulesView, - PermissionCheckView, - JumpView, - InstanceSchemaView, - DatabaseSchemaView, - TableSchemaView, -) -from .views.table import ( - TableAutocompleteView, - TableInsertView, - TableUpsertView, - TableSetColumnTypeView, - TableDropView, - TableFragmentView, - table_view, -) -from .views.row import RowView, RowDeleteView, RowUpdateView -from .renderer import json_renderer -from .url_builder import Urls +from .column_types import SQLiteType +from .csrf import CrossOriginProtectionMiddleware from .database import Database, QueryInterrupted - +from .events import Event +from .plugins import DEFAULT_PLUGINS, get_plugins, pm +from .renderer import json_renderer +from .resources import DatabaseResource, TableResource +from .tokens import TokenInvalid +from .tracer import AsgiTracer +from .url_builder import Urls from .utils import ( + SPATIALITE_FUNCTIONS, PaginatedResources, PrefixedUrlString, - SPATIALITE_FUNCTIONS, StartupError, + add_cors_headers, async_call_with_supported_arguments, await_me_maybe, baseconv, call_with_supported_arguments, detect_json1, - add_cors_headers, display_actor, escape_css_string, escape_sqlite, @@ -121,47 +72,97 @@ from .utils import ( move_plugins_and_allow, move_table_config, parse_metadata, + redact_keys, resolve_env_secrets, resolve_routes, + row_sql_params_pks, sha256_file, tilde_decode, tilde_encode, to_css_class, urlsafe_components, - redact_keys, - row_sql_params_pks, ) -from .tokens import TokenInvalid from .utils.asgi import ( AsgiLifespan, + AsgiRunOnFirstRequest, BadRequest, + DatabaseNotFound, Forbidden, NotFound, - DatabaseNotFound, - TableNotFound, - RowNotFound, Request, Response, - AsgiRunOnFirstRequest, - asgi_static, + RowNotFound, + TableNotFound, asgi_send, asgi_send_file, asgi_send_redirect, + asgi_static, ) -from .csrf import CrossOriginProtectionMiddleware from .utils.internal_db import init_internal_db, populate_schema_tables from .utils.sqlite import ( sqlite3, using_pysqlite3, ) -from .tracer import AsgiTracer -from .plugins import pm, DEFAULT_PLUGINS, get_plugins from .version import __version__ - -from .resources import DatabaseResource, TableResource +from .views import Context +from .views.database import ( + DatabaseView, + QueryView, + database_download, +) +from .views.execute_write import ExecuteWriteAnalyzeView, ExecuteWriteView +from .views.index import IndexView +from .views.row import RowDeleteView, RowUpdateView, RowView +from .views.special import ( + AllowDebugView, + AllowedResourcesView, + ApiExplorerView, + AuthTokenView, + AutocompleteDebugView, + CreateTokenView, + DatabaseSchemaView, + InstanceSchemaView, + JsonDataView, + JumpView, + LogoutView, + MessagesDebugView, + PatternPortfolioView, + PermissionCheckView, + PermissionRulesView, + PermissionsDebugView, + TableSchemaView, +) +from .views.stored_queries import ( + GlobalQueryListView, + QueryCreateAnalyzeView, + QueryDefinitionView, + QueryDeleteView, + QueryEditView, + QueryListView, + QueryParametersView, + QueryStoreView, + QueryUpdateView, +) +from .views.table import ( + TableAutocompleteView, + TableDropView, + TableFragmentView, + TableInsertView, + TableSetColumnTypeView, + TableUpsertView, + table_view, +) +from .views.table_create_alter import ( + DatabaseForeignKeyTargetsView, + TableAlterView, + TableCreateView, + TableForeignKeySuggestionsView, +) app_root = Path(__file__).parent.parent +logger = logging.getLogger(__name__) + # Context variable to track when code is executing within a datasette.client request _in_datasette_client = contextvars.ContextVar("in_datasette_client", default=False) @@ -184,7 +185,7 @@ class PermissionCheck: """Represents a logged permission check for debugging purposes.""" when: str - actor: Dict[str, Any] | None + actor: dict[str, Any] | None action: str parent: str | None child: str | None @@ -434,7 +435,7 @@ class Datasette: if config_dir: db_files = [] for ext in ("db", "sqlite", "sqlite3"): - db_files.extend(config_dir.glob("*.{}".format(ext))) + db_files.extend(config_dir.glob(f"*.{ext}")) self.files += tuple(str(f) for f in db_files) if ( config_dir @@ -675,10 +676,10 @@ class Datasette: def get_jinja_environment(self, request: Request = None) -> Environment: environment = self._jinja_env if request: - for environment in pm.hook.jinja2_environment_from_request( + for hook_environment in pm.hook.jinja2_environment_from_request( datasette=self, request=request, env=environment ): - pass + environment = hook_environment return environment def get_action(self, name_or_abbr: str): @@ -732,7 +733,7 @@ class Datasette: catalog_database_names.update( row["database_name"] for row in await internal_db.execute( - "select distinct database_name from {}".format(table) + f"select distinct database_name from {table}" ) if row["database_name"] is not None ) @@ -743,7 +744,7 @@ class Datasette: for stale_db_name in stale_databases: for table in catalog_table_names: conn.execute( - "DELETE FROM {} WHERE database_name = ?".format(table), + f"DELETE FROM {table} WHERE database_name = ?", [stale_db_name], ) @@ -753,19 +754,7 @@ class Datasette: # Compare schema versions to see if we should skip it if schema_version == current_schema_versions.get(database_name): continue - placeholders = "(?, ?, ?, ?)" - values = [database_name, str(db.path), db.is_memory, schema_version] - if db.path is None: - placeholders = "(?, null, ?, ?)" - values = [database_name, db.is_memory, schema_version] - await internal_db.execute_write( - """ - INSERT OR REPLACE INTO catalog_databases (database_name, path, is_memory, schema_version) - VALUES {} - """.format(placeholders), - values, - ) - await populate_schema_tables(internal_db, db) + await populate_schema_tables(internal_db, db, schema_version) @property def urls(self): @@ -804,17 +793,13 @@ class Datasette: action.name in action_names and action != action_names[action.name] ): - raise StartupError( - "Duplicate action name: {}".format(action.name) - ) + raise StartupError(f"Duplicate action name: {action.name}") if ( action.abbr and action.abbr in action_abbrs and action != action_abbrs[action.abbr] ): - raise StartupError( - "Duplicate action abbr: {}".format(action.abbr) - ) + raise StartupError(f"Duplicate action abbr: {action.abbr}") action_names[action.name] = action if action.abbr: action_abbrs[action.abbr] = action @@ -873,7 +858,7 @@ class Datasette: actor_id: str, *, expires_after: int | None = None, - restrictions: "TokenRestrictions | None" = None, + restrictions: TokenRestrictions | None = None, handler: str | None = None, ) -> str: """ @@ -930,7 +915,7 @@ class Datasette: raise KeyError return matches[0] if name is None: - name = [key for key in self.databases.keys()][0] + name = next(iter(self.databases.keys())) return self.databases[name] def add_database(self, db, name=None, route=None): @@ -943,7 +928,7 @@ class Datasette: suggestion = name i = 2 while name in self.databases: - name = "{}_{}".format(suggestion, i) + name = f"{suggestion}_{i}" i += 1 db.name = name db.route = route or name @@ -978,13 +963,14 @@ class Datasette: for db in dbs: try: db.close() - except Exception as e: + except Exception as e: # noqa: BLE001 + # Collect the first failure and re-raise after every close() has run if first_exception is None: first_exception = e if self.executor is not None: try: self.executor.shutdown(wait=True, cancel_futures=True) - except Exception as e: + except Exception as e: # noqa: BLE001 if first_exception is None: first_exception = e if first_exception is not None: @@ -1333,24 +1319,15 @@ class Datasette: actual = ( actual_sqlite_type.value if actual_sqlite_type is not None - else "unrecognized {!r}".format(column_detail.type) + else f"unrecognized {column_detail.type!r}" ) raise ValueError( - "Column type {!r} is only applicable to SQLite types {} but {}.{}.{} " - "has SQLite type {}".format( - ct_cls.name, - allowed, - database, - resource, - column, - actual, - ) + f"Column type {ct_cls.name!r} is only applicable to SQLite types {allowed} but {database}.{resource}.{column} " + f"has SQLite type {actual}" ) async def _apply_column_types_config(self): """Load column_types from datasette.json config into the internal DB.""" - import logging - for db_name, db_conf in (self.config or {}).get("databases", {}).items(): for table_name, table_conf in db_conf.get("tables", {}).items(): for col_name, ct in table_conf.get("column_types", {}).items(): @@ -1360,7 +1337,7 @@ class Datasette: col_type = ct["type"] config = ct.get("config") if col_type not in self._column_types: - logging.warning( + logger.warning( "column_types config references unknown type %r " "for %s.%s.%s", col_type, @@ -1373,7 +1350,7 @@ class Datasette: db_name, table_name, col_name, col_type, config ) except ValueError as ex: - logging.warning(str(ex)) + logger.warning(str(ex)) async def get_column_type(self, database: str, resource: str, column: str): """ @@ -1426,7 +1403,7 @@ class Datasette: resource: str, column: str, column_type: str, - config: dict = None, + config: dict | None = None, ) -> None: """Assign a column type. Overwrites any existing assignment.""" ct_cls = self._column_types.get(column_type) @@ -1510,9 +1487,7 @@ class Datasette: possible_names = {plugin["name"], plugin["name"].replace("-", "_")} if plugin_name in possible_names: return _resolve_static_asset_path(plugin["static_path"], path) - raise FileNotFoundError( - "No static assets found for plugin {}".format(plugin_name) - ) + raise FileNotFoundError(f"No static assets found for plugin {plugin_name}") def _static_mounted_asset(self, mount_name, path): mount_name = mount_name.strip("/") @@ -1522,7 +1497,7 @@ class Datasette: _resolve_static_asset_path(dirname, path), self.urls.path("/{}/{}".format(mount_name, path.lstrip("/"))), ) - raise FileNotFoundError("No static mount found for {}".format(mount_name)) + raise FileNotFoundError(f"No static mount found for {mount_name}") def _static_asset_hash(self, filepath): filepath = Path(filepath) @@ -1613,18 +1588,17 @@ class Datasette: if await self.allowed(action="view-instance", actor=actor): crumbs.append({"href": self.urls.instance(), "label": "home"}) # Database link - if database: - if await self.allowed( - action="view-database", - resource=DatabaseResource(database=database), - actor=actor, - ): - crumbs.append( - { - "href": self.urls.database(database), - "label": database, - } - ) + if database and await self.allowed( + action="view-database", + resource=DatabaseResource(database=database), + actor=actor, + ): + crumbs.append( + { + "href": self.urls.database(database), + "label": database, + } + ) # Table link if table: assert database, "table= requires database=" @@ -1643,7 +1617,7 @@ class Datasette: async def actors_from_ids( self, actor_ids: Iterable[str | int] - ) -> Dict[int | str, Dict]: + ) -> dict[int | str, dict]: result = pm.hook.actors_from_ids(datasette=self, actor_ids=actor_ids) if result is None: # Do the default thing @@ -1652,9 +1626,9 @@ class Datasette: return result async def track_event(self, event: Event): - assert isinstance(event, self.event_classes), "Invalid event type: {}".format( - type(event) - ) + assert isinstance( + event, self.event_classes + ), f"Invalid event type: {type(event)}" for hook in pm.hook.track_event(datasette=self, event=event): await await_me_maybe(hook) @@ -1691,7 +1665,7 @@ class Datasette: self, actor: dict, action: str, - resource: "Resource" | None = None, + resource: Resource | None = None, ): """ Check if actor can see a resource and if it's private. @@ -1890,10 +1864,7 @@ class Datasette: if truncated and resources: last_resource = resources[-1] # Use tilde-encoding like table pagination - next_token = "{},{}".format( - tilde_encode(str(last_resource.parent)), - tilde_encode(str(last_resource.child)), - ) + next_token = f"{tilde_encode(str(last_resource.parent))},{tilde_encode(str(last_resource.child))}" return PaginatedResources( resources=resources, @@ -1911,7 +1882,7 @@ class Datasette: self, *, action: str, - resource: "Resource" = None, + resource: Resource = None, actor: dict | None = None, ) -> bool: """ @@ -1942,7 +1913,7 @@ class Datasette: self, *, actions: Sequence[str], - resource: "Resource" = None, + resource: Resource = None, actor: dict | None = None, ) -> dict[str, bool]: """ @@ -1963,11 +1934,11 @@ class Datasette: ) # {"edit-schema": True, "drop-table": True, "insert-row": False} """ - from datasette.utils.actions_sql import check_permissions_for_actions from datasette.permissions import ( _permission_check_cache, _skip_permission_checks, ) + from datasette.utils.actions_sql import check_permissions_for_actions # For global actions, resource is None parent = resource.parent if resource else None @@ -2056,7 +2027,7 @@ class Datasette: self, *, action: str, - resource: "Resource" = None, + resource: Resource = None, actor: dict | None = None, ): """ @@ -2110,13 +2081,15 @@ class Datasette: db = self.databases[database] foreign_keys = await db.foreign_keys_for_table(table) # Find the foreign_key for this column - try: - fk = [ + fk = next( + ( foreign_key for foreign_key in foreign_keys if foreign_key["column"] == column - ][0] - except IndexError: + ), + None, + ) + if fk is None: return {} # Ensure user has permission to view the referenced table from datasette.resources import TableResource @@ -2204,16 +2177,17 @@ class Datasette: sqlite_extensions[extension] = result.fetchone()[0] else: sqlite_extensions[extension] = None - except Exception: + except Exception: # noqa: BLE001, S110 + # Probing for optional SQLite extensions - absence is the normal case pass # More details on SpatiaLite if "spatialite" in sqlite_extensions: spatialite_details = {} for fn in SPATIALITE_FUNCTIONS: try: - result = conn.execute("select {}()".format(fn)) + result = conn.execute(f"select {fn}()") spatialite_details[fn] = result.fetchone()[0] - except Exception as e: + except sqlite3.Error as e: spatialite_details[fn] = {"error": str(e)} sqlite_extensions["spatialite"] = spatialite_details @@ -2221,9 +2195,7 @@ class Datasette: fts_versions = [] for fts in ("FTS5", "FTS4", "FTS3"): try: - conn.execute( - "CREATE VIRTUAL TABLE v{fts} USING {fts} (data)".format(fts=fts) - ) + conn.execute(f"CREATE VIRTUAL TABLE v{fts} USING {fts} (data)") fts_versions.append(fts) except sqlite3.OperationalError: continue @@ -2282,7 +2254,7 @@ class Datasette: "static": p["static_path"] is not None, "templates": p["templates_path"] is not None, "version": p.get("version"), - "hooks": list(sorted(set(p["hooks"]))), + "hooks": sorted(set(p["hooks"])), } for p in ps ] @@ -2358,13 +2330,15 @@ class Datasette: async def render_template( self, - templates: List[str] | str | Template, - context: Dict[str, Any] | Context | None = None, + templates: list[str] | str | Template, + context: dict[str, Any] | Context | None = None, request: Request | None = None, view_name: str | None = None, ): if not self._startup_invoked: - raise Exception("render_template() called before await ds.invoke_startup()") + raise RuntimeError( + "render_template() called before await ds.invoke_startup()" + ) context = context or {} if isinstance(templates, Template): template = templates @@ -2410,9 +2384,9 @@ class Datasette: datasette=self, ): extra_vars = await await_me_maybe(extra_vars) - assert isinstance(extra_vars, dict), "extra_vars is of type {}".format( - type(extra_vars) - ) + assert isinstance( + extra_vars, dict + ), f"extra_vars is of type {type(extra_vars)}" extra_template_vars.update(extra_vars) async def menu_links(): @@ -2431,29 +2405,27 @@ class Datasette: # the contract tests fail otherwise template_context = { **context, - **{ - "request": request, - "crumb_items": self._crumb_items, - "urls": self.urls, - "actor": request.actor if request else None, - "menu_links": menu_links, - "display_actor": display_actor, - "show_logout": request is not None - and "ds_actor" in request.cookies - and request.actor, - "zip": zip, - "body_scripts": body_scripts, - "format_bytes": format_bytes, - "show_messages": lambda: self._show_messages(request), - "extra_css_urls": await self._asset_urls( - "extra_css_urls", template, context, request, view_name - ), - "extra_js_urls": await self._asset_urls( - "extra_js_urls", template, context, request, view_name - ), - "base_url": self.setting("base_url"), - "datasette_version": __version__, - }, + "request": request, + "crumb_items": self._crumb_items, + "urls": self.urls, + "actor": request.actor if request else None, + "menu_links": menu_links, + "display_actor": display_actor, + "show_logout": request is not None + and "ds_actor" in request.cookies + and request.actor, + "zip": zip, + "body_scripts": body_scripts, + "format_bytes": format_bytes, + "show_messages": lambda: self._show_messages(request), + "extra_css_urls": await self._asset_urls( + "extra_css_urls", template, context, request, view_name + ), + "extra_js_urls": await self._asset_urls( + "extra_js_urls", template, context, request, view_name + ), + "base_url": self.setting("base_url"), + "datasette_version": __version__, **extra_template_vars, } if request and request.args.get("_context") and self.setting("template_debug"): @@ -2575,7 +2547,7 @@ class Datasette: JsonDataView.as_view( self, "plugins.json", - lambda request: {"plugins": self._plugins(request)}, + self._plugins, needs_request=True, ), r"/-/plugins(\.(?Pjson))?$", @@ -2953,7 +2925,8 @@ class DatasetteRouter: custom_response ), "Default forbidden() hook should have been called" return await custom_response.asgi_send(send) - except Exception as exception: + except Exception as exception: # noqa: BLE001 + # This IS the top-level error handler - it must catch everything return await self.handle_exception(request, send, exception) async def handle_401(self, request, send, exception): @@ -2975,7 +2948,7 @@ class DatasetteRouter: request.path.replace("~", "~7E").replace("%", "~").replace(".", "~2E") ) if request.query_string: - new_path += "?{}".format(request.query_string) + new_path += f"?{request.query_string}" await asgi_send_redirect(send, new_path) return # If URL has a trailing slash, redirect to URL without it @@ -3185,8 +3158,7 @@ _curly_re = re.compile(r"({.*?})") def route_pattern_from_filepath(filepath): # Drop the ".html" suffix - if filepath.endswith(".html"): - filepath = filepath[: -len(".html")] + filepath = filepath.removesuffix(".html") re_bits = ["/"] for bit in _curly_re.split(filepath): if _curly_re.match(bit): diff --git a/datasette/blob_renderer.py b/datasette/blob_renderer.py index 4d8c6bea..b6c8b77f 100644 --- a/datasette/blob_renderer.py +++ b/datasette/blob_renderer.py @@ -1,8 +1,9 @@ -from datasette import hookimpl -from datasette.utils.asgi import Response, BadRequest -from datasette.utils import to_css_class import hashlib +from datasette import hookimpl +from datasette.utils import to_css_class +from datasette.utils.asgi import BadRequest, Response + _BLOB_COLUMN = "_blob_column" _BLOB_HASH = "_blob_hash" diff --git a/datasette/cli.py b/datasette/cli.py index 90a33e80..57db83b6 100644 --- a/datasette/cli.py +++ b/datasette/cli.py @@ -1,43 +1,45 @@ import asyncio -import uvicorn -import click -from click import formatting -from click.types import CompositeParamType -from click_default_group import DefaultGroup import functools import json import os import pathlib -from runpy import run_module import shutil -from subprocess import call import sys import textwrap import webbrowser +from runpy import run_module +from subprocess import call + +import click +import uvicorn +from click import formatting +from click.types import CompositeParamType +from click_default_group import DefaultGroup + from .app import ( - Datasette, DEFAULT_SETTINGS, SETTINGS, SQLITE_LIMIT_ATTACHED, + Datasette, pm, ) from .inspect import inspect_tables from .utils import ( + ConnectionProblem, LoadExtension, + SpatialiteConnectionProblem, + SpatialiteNotFound, StartupError, + StaticMount, + ValueAsBooleanError, check_connection, deep_dict_update, find_spatialite, - parse_metadata, - ConnectionProblem, - SpatialiteConnectionProblem, initial_path_for_datasette, pairs_to_nested_config, + parse_metadata, temporary_docker_directory, value_as_boolean, - SpatialiteNotFound, - StaticMount, - ValueAsBooleanError, ) from .utils.sqlite import sqlite3 from .utils.testing import TestClient @@ -75,7 +77,7 @@ class Setting(CompositeParamType): # Datasette 1.0, we turn bare setting names into setting.name # Type checking for those older settings default = DEFAULT_SETTINGS[name] - name = "settings.{}".format(name) + name = f"settings.{name}" if isinstance(default, bool): try: return name, "true" if value_as_boolean(value) else "false" @@ -171,7 +173,6 @@ async def inspect_(files, sqlite_extensions): @cli.group() def publish(): """Publish specified SQLite database files to the internet along with a Datasette-powered interface and API""" - pass # Register publish plugins @@ -578,27 +579,27 @@ def serve( # https://github.com/simonw/datasette/issues/2389 deep_dict_update(config_data, settings_updates) - kwargs = dict( - immutables=immutable, - cache_headers=not reload, - cors=cors, - inspect_data=inspect_data, - config=config_data, - metadata=metadata_data, - sqlite_extensions=sqlite_extensions, - template_dir=template_dir, - plugins_dir=plugins_dir, - static_mounts=static, - settings=None, # These are passed in config= now - memory=memory, - secret=secret, - version_note=version_note, - pdb=pdb, - crossdb=crossdb, - nolock=nolock, - internal=internal, - default_deny=default_deny, - ) + kwargs = { + "immutables": immutable, + "cache_headers": not reload, + "cors": cors, + "inspect_data": inspect_data, + "config": config_data, + "metadata": metadata_data, + "sqlite_extensions": sqlite_extensions, + "template_dir": template_dir, + "plugins_dir": plugins_dir, + "static_mounts": static, + "settings": None, # These are passed in config= now + "memory": memory, + "secret": secret, + "version_note": version_note, + "pdb": pdb, + "crossdb": crossdb, + "nolock": nolock, + "internal": internal, + "default_deny": default_deny, + } # Separate directories from files directories = [f for f in files if os.path.isdir(f)] @@ -621,9 +622,7 @@ def serve( conn.close() else: raise click.ClickException( - "Invalid value for '[FILES]...': Path '{}' does not exist.".format( - file - ) + f"Invalid value for '[FILES]...': Path '{file}' does not exist." ) # Check for duplicate files by resolving all paths to their absolute forms @@ -684,7 +683,7 @@ def serve( client = TestClient(ds) request_headers = {} if token: - request_headers["Authorization"] = "Bearer {}".format(token) + request_headers["Authorization"] = f"Bearer {token}" cookies = {} if actor: cookies["ds_actor"] = client.actor_cookie(json.loads(actor)) @@ -719,9 +718,13 @@ def serve( path = run_sync(lambda: initial_path_for_datasette(ds)) url = f"http://{host}:{port}{path}" webbrowser.open(url) - uvicorn_kwargs = dict( - host=host, port=port, log_level="info", lifespan="on", workers=1 - ) + uvicorn_kwargs = { + "host": host, + "port": port, + "log_level": "info", + "lifespan": "on", + "workers": 1, + } if uds: uvicorn_kwargs["uds"] = uds if ssl_keyfile: @@ -885,7 +888,7 @@ async def check_databases(ds): ) except ConnectionProblem as e: raise click.UsageError( - f"Connection to {database.path} failed check: {str(e.args[0])}" + f"Connection to {database.path} failed check: {e.args[0]!s}" ) # If --crossdb and more than SQLITE_LIMIT_ATTACHED show warning if ( @@ -893,9 +896,5 @@ async def check_databases(ds): and len([db for db in ds.databases.values() if not db.is_memory]) > SQLITE_LIMIT_ATTACHED ): - msg = ( - "Warning: --crossdb only works with the first {} attached databases".format( - SQLITE_LIMIT_ATTACHED - ) - ) + msg = f"Warning: --crossdb only works with the first {SQLITE_LIMIT_ATTACHED} attached databases" click.echo(click.style(msg, bold=True, fg="yellow"), err=True) diff --git a/datasette/column_types.py b/datasette/column_types.py index 11a14ec0..92fdd969 100644 --- a/datasette/column_types.py +++ b/datasette/column_types.py @@ -64,14 +64,14 @@ class ColumnType: Return an HTML string to render this cell value, or None to fall through to the default render_cell plugin hook chain. """ - return None + return async def validate(self, value, datasette): """ Validate a value before it is written. Return None if valid, or a string error message if invalid. """ - return None + return async def transform_value(self, value, datasette): """ diff --git a/datasette/csrf.py b/datasette/csrf.py index df239aee..a62f9473 100644 --- a/datasette/csrf.py +++ b/datasette/csrf.py @@ -40,12 +40,12 @@ def _origin_tuple(value): scheme = (parsed.scheme or "").lower() host = (parsed.hostname or "").lower() if not scheme or not host: - raise ValueError("missing scheme or host in {!r}".format(value)) + raise ValueError(f"missing scheme or host in {value!r}") port = parsed.port # may raise ValueError on bad ports if port is None: port = DEFAULT_PORTS.get(scheme) if port is None: - raise ValueError("unknown default port for scheme {!r}".format(scheme)) + raise ValueError(f"unknown default port for scheme {scheme!r}") return scheme, host, port @@ -125,9 +125,7 @@ class CrossOriginProtectionMiddleware: return await self._forbid( send, - "Sec-Fetch-Site was {!r}, expected 'same-origin' or 'none'".format( - sec_fetch_site - ), + f"Sec-Fetch-Site was {sec_fetch_site!r}, expected 'same-origin' or 'none'", ) return @@ -141,11 +139,11 @@ class CrossOriginProtectionMiddleware: request_scheme = self._request_scheme(scope) try: origin_tuple = _origin_tuple(origin) - expected_tuple = _origin_tuple("{}://{}".format(request_scheme, host)) + expected_tuple = _origin_tuple(f"{request_scheme}://{host}") except ValueError: await self._forbid( send, - "Malformed Origin {!r} or Host {!r}".format(origin, host), + f"Malformed Origin {origin!r} or Host {host!r}", ) return @@ -155,7 +153,7 @@ class CrossOriginProtectionMiddleware: await self._forbid( send, - "Origin {!r} does not match Host {!r}".format(origin, host), + f"Origin {origin!r} does not match Host {host!r}", ) def _request_scheme(self, scope): @@ -163,7 +161,8 @@ class CrossOriginProtectionMiddleware: try: if self.datasette.setting("force_https_urls"): return "https" - except Exception: + except Exception: # noqa: BLE001, S110 + # Settings may not be readable this early; fall back to the ASGI scheme pass return scope.get("scheme") or "http" diff --git a/datasette/database.py b/datasette/database.py index e7fe1ed9..e162d34e 100644 --- a/datasette/database.py +++ b/datasette/database.py @@ -1,33 +1,35 @@ import asyncio import atexit -from collections import namedtuple import inspect import os -from pathlib import Path import queue -import sqlite_utils import sys import tempfile import threading import uuid +from collections import namedtuple +from pathlib import Path +import sqlite_utils + +from .inspect import inspect_hash from .tracer import trace from .utils import ( call_with_supported_arguments, detect_fts, detect_primary_keys, detect_spatialite, + escape_sqlite, get_all_foreign_keys, get_outbound_foreign_keys, md5_not_usedforsecurity, - sqlite_timelimit, sqlite3, - table_columns, + sqlite_timelimit, table_column_details, + table_columns, ) from .utils.sql_analysis import SQLAnalysis, analyze_sql_tables from .utils.sqlite import sqlite_hidden_table_names -from .inspect import inspect_hash connections = threading.local() @@ -98,9 +100,7 @@ class Database: def _check_not_closed(self): if self._closed: - raise DatasetteClosedError( - "Database {!r} has been closed".format(self.name) - ) + raise DatasetteClosedError(f"Database {self.name!r} has been closed") def _remove_pending_execute_future(self, future): with self._pending_execute_futures_lock: @@ -139,7 +139,7 @@ class Database: if write: extra_kwargs["isolation_level"] = "IMMEDIATE" if self.memory_name: - uri = "file:{}?mode=memory&cache=shared".format(self.memory_name) + uri = f"file:{self.memory_name}?mode=memory&cache=shared" conn = sqlite3.connect( uri, uri=True, check_same_thread=False, **extra_kwargs ) @@ -192,21 +192,20 @@ class Database: write_thread.join(timeout=10) if write_thread.is_alive(): sys.stderr.write( - "Datasette: write thread for {!r} did not exit within 10s\n".format( - self.name - ) + f"Datasette: write thread for {self.name!r} did not exit within 10s\n" ) sys.stderr.flush() for future in pending_execute_futures: try: future.result() - except Exception: + except Exception: # noqa: BLE001, S110 + # Shutdown teardown - a failed pending write must not block close() pass # Close anything still tracked in _all_file_connections for connection in self._all_file_connections: try: connection.close() - except Exception: + except Exception: # noqa: BLE001, S110 pass self._all_file_connections = [] # Drop per-thread cached read connections we can reach @@ -218,13 +217,13 @@ class Database: if self._read_connection is not None: try: self._read_connection.close() - except Exception: + except Exception: # noqa: BLE001, S110 pass self._read_connection = None if self._write_connection is not None: try: self._write_connection.close() - except Exception: + except Exception: # noqa: BLE001, S110 pass self._write_connection = None if self.is_temp_disk: @@ -246,6 +245,7 @@ class Database: request=None, return_all=False, returning_limit=EXECUTE_WRITE_RETURNING_LIMIT, + transaction=True, ): self._check_not_closed() if returning_limit < 0: @@ -258,7 +258,9 @@ class Database: ) with trace("sql", database=self.name, sql=sql.strip(), params=params): - results = await self.execute_write_fn(_inner, block=block, request=request) + results = await self.execute_write_fn( + _inner, block=block, request=request, transaction=transaction + ) return results async def execute_write_script(self, sql, block=True, request=None): @@ -348,6 +350,7 @@ class Database: self.ds._prepare_connection(self._write_connection, self.name) if transaction: with self._write_connection: + self._write_connection.execute("BEGIN IMMEDIATE") result = fn(self._write_connection) else: result = fn(self._write_connection) @@ -366,7 +369,8 @@ class Database: async def _dispatch_events_after_write(): try: await reply_future - except Exception: + except Exception: # noqa: BLE001 + # The write failed; skip success events regardless of why # if the write failed, don't emit success events return for event in pending_events: @@ -419,9 +423,7 @@ class Database: self._write_thread = threading.Thread( target=self._execute_writes, daemon=True ) - self._write_thread.name = "_execute_writes for database {}".format( - self.name - ) + self._write_thread.name = f"_execute_writes for database {self.name}" self._write_thread.start() task_id = uuid.uuid5(uuid.NAMESPACE_DNS, "datasette.io") loop = asyncio.get_running_loop() @@ -442,7 +444,8 @@ class Database: try: conn = self.connect(write=True) self.ds._prepare_connection(conn, self.name) - except Exception as e: + except Exception as e: # noqa: BLE001 + # Stored and re-raised to whoever queues the next write conn_exception = e while True: task = self._write_queue.get() @@ -450,7 +453,8 @@ class Database: if conn is not None: try: conn.close() - except Exception: + except Exception: # noqa: BLE001, S110 + # Best-effort close as the write thread exits pass return exception = None @@ -469,19 +473,21 @@ class Database: except ValueError: # Was probably a memory connection pass - except Exception as e: - sys.stderr.write("{}\n".format(e)) + except Exception as e: # noqa: BLE001 + # Write thread must survive any task failure or the database wedges + sys.stderr.write(f"{e}\n") sys.stderr.flush() exception = e else: try: if task.transaction: with conn: + conn.execute("BEGIN IMMEDIATE") result = task.fn(conn) else: result = task.fn(conn) - except Exception as e: - sys.stderr.write("{}\n".format(e)) + except Exception as e: # noqa: BLE001 + sys.stderr.write(f"{e}\n") sys.stderr.flush() exception = e _deliver_write_result(task, result, exception) @@ -548,9 +554,7 @@ class Database: raise QueryInterrupted(e, sql, params) if log_sql_errors: sys.stderr.write( - "ERROR: conn={}, sql = {}, params = {}: {}\n".format( - conn, repr(sql), params, e - ) + f"ERROR: conn={conn}, sql = {sql!r}, params = {params}: {e}\n" ) sys.stderr.flush() raise @@ -603,7 +607,7 @@ class Database: try: table_count = ( await self.execute( - f"select count(*) from (select * from [{table}] limit {self.count_limit + 1})", + f"select count(*) from (select * from {escape_sqlite(table)} limit {self.count_limit + 1})", custom_time_limit=limit, ) ).rows[0][0] @@ -707,9 +711,9 @@ class Database: column_names and len(column_names) == 2 and ("id" in column_names or "pk" in column_names) - and not set(column_names) == {"id", "pk"} + and set(column_names) != {"id", "pk"} ): - return [c for c in column_names if c not in ("id", "pk")][0] + return next(c for c in column_names if c not in ("id", "pk")) # Couldn't find a label: return None @@ -851,10 +855,10 @@ def _apply_write_wrapper(fn, wrapper_factory, track_event): class WriteTask: __slots__ = ( "fn", - "task_id", + "isolated_connection", "loop", "reply_future", - "isolated_connection", + "task_id", "transaction", ) @@ -895,7 +899,7 @@ class QueryInterrupted(Exception): self.params = params def __str__(self): - return "QueryInterrupted: {}".format(self.e) + return f"QueryInterrupted: {self.e}" class MultipleValues(Exception): diff --git a/datasette/default_actions.py b/datasette/default_actions.py index 602e0df4..ee165ae5 100644 --- a/datasette/default_actions.py +++ b/datasette/default_actions.py @@ -2,8 +2,8 @@ from datasette import hookimpl from datasette.permissions import Action from datasette.resources import ( DatabaseResource, - TableResource, QueryResource, + TableResource, ) diff --git a/datasette/default_magic_parameters.py b/datasette/default_magic_parameters.py index 91c1c5aa..bff5f0a7 100644 --- a/datasette/default_magic_parameters.py +++ b/datasette/default_magic_parameters.py @@ -1,8 +1,9 @@ -from datasette import hookimpl import datetime import os import time +from datasette import hookimpl + def header(key, request): key = key.replace("_", "-").encode("utf-8") diff --git a/datasette/default_permissions/__init__.py b/datasette/default_permissions/__init__.py index 6cd46f04..dee5df42 100644 --- a/datasette/default_permissions/__init__.py +++ b/datasette/default_permissions/__init__.py @@ -17,18 +17,29 @@ UNION/INTERSECT operations. The order of evaluation is: from __future__ import annotations -# Re-export all hooks and public utilities -from .restrictions import ( - actor_restrictions_sql as actor_restrictions_sql, - restrictions_allow_action as restrictions_allow_action, - ActorRestrictions as ActorRestrictions, -) -from .root import root_user_permissions_sql as root_user_permissions_sql from .config import config_permissions_sql as config_permissions_sql +from .defaults import ( + DEFAULT_ALLOW_ACTIONS as DEFAULT_ALLOW_ACTIONS, +) +from .defaults import ( + default_action_permissions_sql as default_action_permissions_sql, +) from .defaults import ( # Avoid "datasette.default_permissions" does not explicitly export attribute default_allow_sql_check as default_allow_sql_check, - default_action_permissions_sql as default_action_permissions_sql, - default_query_permissions_sql as default_query_permissions_sql, - DEFAULT_ALLOW_ACTIONS as DEFAULT_ALLOW_ACTIONS, ) +from .defaults import ( + default_query_permissions_sql as default_query_permissions_sql, +) +from .restrictions import ( + ActorRestrictions as ActorRestrictions, +) + +# Re-export all hooks and public utilities +from .restrictions import ( + actor_restrictions_sql as actor_restrictions_sql, +) +from .restrictions import ( + restrictions_allow_action as restrictions_allow_action, +) +from .root import root_user_permissions_sql as root_user_permissions_sql diff --git a/datasette/default_permissions/config.py b/datasette/default_permissions/config.py index aab87c1c..4494f07f 100644 --- a/datasette/default_permissions/config.py +++ b/datasette/default_permissions/config.py @@ -6,7 +6,7 @@ Applies permission rules from datasette.yaml configuration. from __future__ import annotations -from typing import TYPE_CHECKING, Any, List, Optional, Set, Tuple +from typing import TYPE_CHECKING, Any if TYPE_CHECKING: from datasette.app import Datasette @@ -55,8 +55,8 @@ class ConfigPermissionProcessor: def __init__( self, - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, ): self.datasette = datasette @@ -74,8 +74,8 @@ class ConfigPermissionProcessor: self.restrictions = actor.get("_r", {}) if actor else {} # Pre-compute restriction info for efficiency - self.restricted_databases: Set[str] = set() - self.restricted_tables: Set[Tuple[str, str]] = set() + self.restricted_databases: set[str] = set() + self.restricted_tables: set[tuple[str, str]] = set() if self.has_restrictions: self.restricted_databases = { @@ -92,16 +92,20 @@ class ConfigPermissionProcessor: # Tables implicitly reference their parent databases self.restricted_databases.update(db for db, _ in self.restricted_tables) - def evaluate_allow_block(self, allow_block: Any) -> Optional[bool]: + def evaluate_allow_block(self, allow_block: Any) -> bool | None: """Evaluate an allow block against the current actor.""" if allow_block is None: return None + # Values passed using ``-s permissions.* 1`` or ``0`` are parsed as + # integers, but should retain the CLI's boolean 1/0 behavior. + if isinstance(allow_block, int) and allow_block in (0, 1): + return bool(allow_block) return actor_matches_allow(self.actor, allow_block) def is_in_restriction_allowlist( self, - parent: Optional[str], - child: Optional[str], + parent: str | None, + child: str | None, ) -> bool: """Check if resource is allowed by actor restrictions.""" if not self.has_restrictions: @@ -143,9 +147,9 @@ class ConfigPermissionProcessor: def add_permissions_rule( self, - parent: Optional[str], - child: Optional[str], - permissions_block: Optional[dict], + parent: str | None, + child: str | None, + permissions_block: dict | None, scope_desc: str, ) -> None: """Add a rule from a permissions:{action} block.""" @@ -165,8 +169,8 @@ class ConfigPermissionProcessor: def add_allow_block_rule( self, - parent: Optional[str], - child: Optional[str], + parent: str | None, + child: str | None, allow_block: Any, scope_desc: str, ) -> None: @@ -198,8 +202,8 @@ class ConfigPermissionProcessor: def _add_restriction_gate_denies( self, - parent: Optional[str], - child: Optional[str], + parent: str | None, + child: str | None, is_allowed: bool, scope_desc: str, ) -> None: @@ -231,7 +235,7 @@ class ConfigPermissionProcessor: if db_name == parent: self.collector.add(db_name, table_name, False, reason) - def process(self) -> Optional[PermissionSQL]: + def process(self) -> PermissionSQL | None: """Process all config rules and return combined PermissionSQL.""" self._process_root_permissions() self._process_databases() @@ -421,10 +425,10 @@ class ConfigPermissionProcessor: @hookimpl(specname="permission_resources_sql") async def config_permissions_sql( - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, -) -> Optional[List[PermissionSQL]]: +) -> list[PermissionSQL] | None: """ Apply permission rules from datasette.yaml configuration. diff --git a/datasette/default_permissions/defaults.py b/datasette/default_permissions/defaults.py index 5bc74425..6f97812b 100644 --- a/datasette/default_permissions/defaults.py +++ b/datasette/default_permissions/defaults.py @@ -6,7 +6,7 @@ Provides default allow rules for standard view/execute actions. from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: from datasette.app import Datasette @@ -29,29 +29,28 @@ DEFAULT_ALLOW_ACTIONS = frozenset( @hookimpl(specname="permission_resources_sql") async def default_allow_sql_check( - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, -) -> Optional[PermissionSQL]: +) -> PermissionSQL | None: """ Enforce the default_allow_sql setting. When default_allow_sql is false (the default), execute-sql is denied unless explicitly allowed by config or other rules. """ - if action == "execute-sql": - if not datasette.setting("default_allow_sql"): - return PermissionSQL.deny(reason="default_allow_sql is false") + if action == "execute-sql" and not datasette.setting("default_allow_sql"): + return PermissionSQL.deny(reason="default_allow_sql is false") return None @hookimpl(specname="permission_resources_sql") async def default_action_permissions_sql( - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, -) -> Optional[PermissionSQL]: +) -> PermissionSQL | None: """ Provide default allow rules for standard view/execute actions. @@ -71,10 +70,10 @@ async def default_action_permissions_sql( @hookimpl(specname="permission_resources_sql") async def default_query_permissions_sql( - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, -) -> Optional[PermissionSQL]: +) -> PermissionSQL | None: actor_id = actor.get("id") if isinstance(actor, dict) else None if action not in {"view-query", "update-query", "delete-query"}: diff --git a/datasette/default_permissions/helpers.py b/datasette/default_permissions/helpers.py index 47e03569..5e59b7b4 100644 --- a/datasette/default_permissions/helpers.py +++ b/datasette/default_permissions/helpers.py @@ -5,7 +5,7 @@ Shared helper utilities for default permission implementations. from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, List, Optional, Set +from typing import TYPE_CHECKING if TYPE_CHECKING: from datasette.app import Datasette @@ -13,7 +13,7 @@ if TYPE_CHECKING: from datasette.permissions import PermissionSQL -def get_action_name_variants(datasette: "Datasette", action: str) -> Set[str]: +def get_action_name_variants(datasette: Datasette, action: str) -> set[str]: """ Get all name variants for an action (full name and abbreviation). @@ -27,7 +27,7 @@ def get_action_name_variants(datasette: "Datasette", action: str) -> Set[str]: return variants -def action_in_list(datasette: "Datasette", action: str, action_list: list) -> bool: +def action_in_list(datasette: Datasette, action: str, action_list: list) -> bool: """Check if an action (or its abbreviation) is in a list.""" return bool(get_action_name_variants(datasette, action).intersection(action_list)) @@ -36,8 +36,8 @@ def action_in_list(datasette: "Datasette", action: str, action_list: list) -> bo class PermissionRow: """A single permission rule row.""" - parent: Optional[str] - child: Optional[str] + parent: str | None + child: str | None allow: bool reason: str @@ -46,14 +46,14 @@ class PermissionRowCollector: """Collects permission rows and converts them to PermissionSQL.""" def __init__(self, prefix: str = "row"): - self.rows: List[PermissionRow] = [] + self.rows: list[PermissionRow] = [] self.prefix = prefix def add( self, - parent: Optional[str], - child: Optional[str], - allow: Optional[bool], + parent: str | None, + child: str | None, + allow: bool | None, reason: str, if_not_none: bool = False, ) -> None: @@ -62,7 +62,7 @@ class PermissionRowCollector: return self.rows.append(PermissionRow(parent, child, allow, reason)) - def to_permission_sql(self) -> Optional[PermissionSQL]: + def to_permission_sql(self) -> PermissionSQL | None: """Convert collected rows to a PermissionSQL object.""" if not self.rows: return None diff --git a/datasette/default_permissions/restrictions.py b/datasette/default_permissions/restrictions.py index a22cd7e5..88e1d274 100644 --- a/datasette/default_permissions/restrictions.py +++ b/datasette/default_permissions/restrictions.py @@ -8,7 +8,7 @@ contains allowlists of resources the actor can access. from __future__ import annotations from dataclasses import dataclass -from typing import TYPE_CHECKING, List, Optional, Set, Tuple +from typing import TYPE_CHECKING if TYPE_CHECKING: from datasette.app import Datasette @@ -23,12 +23,12 @@ from .helpers import action_in_list, get_action_name_variants class ActorRestrictions: """Parsed actor restrictions from the _r key.""" - global_actions: List[str] # _r.a - globally allowed actions + global_actions: list[str] # _r.a - globally allowed actions database_actions: dict # _r.d - {db_name: [actions]} table_actions: dict # _r.r - {db_name: {table: [actions]}} @classmethod - def from_actor(cls, actor: Optional[dict]) -> Optional["ActorRestrictions"]: + def from_actor(cls, actor: dict | None) -> ActorRestrictions | None: """Parse restrictions from actor dict. Returns None if no restrictions.""" if not actor: return None @@ -44,11 +44,11 @@ class ActorRestrictions: table_actions=restrictions.get("r", {}), ) - def is_action_globally_allowed(self, datasette: "Datasette", action: str) -> bool: + def is_action_globally_allowed(self, datasette: Datasette, action: str) -> bool: """Check if action is in the global allowlist.""" return action_in_list(datasette, action, self.global_actions) - def get_allowed_databases(self, datasette: "Datasette", action: str) -> Set[str]: + def get_allowed_databases(self, datasette: Datasette, action: str) -> set[str]: """Get database names where this action is allowed.""" allowed = set() for db_name, db_actions in self.database_actions.items(): @@ -57,8 +57,8 @@ class ActorRestrictions: return allowed def get_allowed_tables( - self, datasette: "Datasette", action: str - ) -> Set[Tuple[str, str]]: + self, datasette: Datasette, action: str + ) -> set[tuple[str, str]]: """Get (database, table) pairs where this action is allowed.""" allowed = set() for db_name, tables in self.table_actions.items(): @@ -70,10 +70,10 @@ class ActorRestrictions: @hookimpl(specname="permission_resources_sql") async def actor_restrictions_sql( - datasette: "Datasette", - actor: Optional[dict], + datasette: Datasette, + actor: dict | None, action: str, -) -> Optional[List[PermissionSQL]]: +) -> list[PermissionSQL] | None: """ Handle actor restriction-based permission rules. @@ -140,10 +140,10 @@ async def actor_restrictions_sql( def restrictions_allow_action( - datasette: "Datasette", + datasette: Datasette, restrictions: dict, action: str, - resource: Optional[str | Tuple[str, str]], + resource: str | tuple[str, str] | None, ) -> bool: """ Check if restrictions allow the requested action on the requested resource. diff --git a/datasette/default_permissions/root.py b/datasette/default_permissions/root.py index 4931f7ff..22d13f65 100644 --- a/datasette/default_permissions/root.py +++ b/datasette/default_permissions/root.py @@ -6,7 +6,7 @@ Grants full permissions to the root user when --root flag is used. from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: from datasette.app import Datasette @@ -17,9 +17,9 @@ from datasette.permissions import PermissionSQL @hookimpl(specname="permission_resources_sql") async def root_user_permissions_sql( - datasette: "Datasette", - actor: Optional[dict], -) -> Optional[PermissionSQL]: + datasette: Datasette, + actor: dict | None, +) -> PermissionSQL | None: """ Grant root user full permissions when --root flag is used. """ diff --git a/datasette/default_permissions/tokens.py b/datasette/default_permissions/tokens.py index 7a359dc6..52daf8a2 100644 --- a/datasette/default_permissions/tokens.py +++ b/datasette/default_permissions/tokens.py @@ -7,7 +7,7 @@ to datasette.verify_token() so all registered handlers are tried. from __future__ import annotations -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING if TYPE_CHECKING: from datasette.app import Datasette @@ -17,15 +17,13 @@ from datasette.tokens import SignedTokenHandler @hookimpl -def register_token_handler(datasette: "Datasette"): +def register_token_handler(datasette: Datasette): """Register the default signed token handler.""" return SignedTokenHandler() @hookimpl(specname="actor_from_request") -async def actor_from_signed_api_token( - datasette: "Datasette", request -) -> Optional[dict]: +async def actor_from_signed_api_token(datasette: Datasette, request) -> dict | None: """ Authenticate requests using API tokens by delegating to all registered token handlers via datasette.verify_token(). diff --git a/datasette/default_table_actions.py b/datasette/default_table_actions.py index e41434ef..0f2f32ef 100644 --- a/datasette/default_table_actions.py +++ b/datasette/default_table_actions.py @@ -20,7 +20,7 @@ def table_actions(datasette, actor, database, table, request): "label": "Alter table", "description": "Change columns and primary key for this table.", "attrs": { - "aria-label": "Alter table {}".format(table), + "aria-label": f"Alter table {table}", "data-table-action": "alter-table", }, } diff --git a/datasette/events.py b/datasette/events.py index e8786da9..5f3fd06e 100644 --- a/datasette/events.py +++ b/datasette/events.py @@ -1,8 +1,9 @@ from abc import ABC, abstractproperty from dataclasses import asdict, dataclass, field -from datasette.hookspecs import hookimpl from datetime import datetime, timezone +from datasette.hookspecs import hookimpl + @dataclass class Event(ABC): diff --git a/datasette/facets.py b/datasette/facets.py index abe0605e..8c09e1dc 100644 --- a/datasette/facets.py +++ b/datasette/facets.py @@ -1,12 +1,13 @@ import json import urllib + from datasette import hookimpl from datasette.database import QueryInterrupted from datasette.utils import ( + detect_json1, escape_sqlite, path_with_added_args, path_with_removed_args, - detect_json1, sqlite3, ) @@ -30,7 +31,7 @@ def load_facet_configs(request, table_config): assert ( len(facet_config.values()) == 1 ), "Metadata config dicts should be {type: config}" - type, facet_config = list(facet_config.items())[0] + type, facet_config = next(iter(facet_config.items())) if isinstance(facet_config, str): facet_config = {"simple": facet_config} facet_configs.setdefault(type, []).append( @@ -85,7 +86,7 @@ class Facet: self.database = database # For foreign key expansion. Can be None for e.g. stored SQL queries: self.table = table - self.sql = sql or f"select * from [{table}]" + self.sql = sql or f"select * from {escape_sqlite(table)}" self.params = params or [] self.table_config = table_config # row_count can be None, in which case we calculate it ourselves: @@ -160,18 +161,13 @@ class ColumnFacet(Facet): for column in columns: if column in already_enabled: continue - suggested_facet_sql = """ - with limited as (select * from ({sql}) limit {suggest_consider}) - select {column} as value, count(*) as n from limited + suggested_facet_sql = f""" + with limited as (select * from ({self.sql}) limit {self.suggest_consider}) + select {escape_sqlite(column)} as value, count(*) as n from limited where value is not null group by value - limit {limit} - """.format( - column=escape_sqlite(column), - sql=self.sql, - limit=facet_size + 1, - suggest_consider=self.suggest_consider, - ) + limit {facet_size + 1} + """ distinct_values = None try: distinct_values = await self.ds.execute( @@ -267,7 +263,7 @@ class ColumnFacet(Facet): for row in facet_rows: column_qs = column if column.startswith("_"): - column_qs = "{}__exact".format(column) + column_qs = f"{column}__exact" selected = (column_qs, str(row["value"])) in qs_pairs if selected: toggle_path = path_with_removed_args( @@ -342,12 +338,12 @@ class ArrayFacet(Facet): for v in await self.ds.execute( self.database, ( - "select {column} from ({sql}) " - "where {column} is not null " - "and {column} != '' " - "and json_array_length({column}) > 0 " + f"select {escape_sqlite(column)} from ({self.sql}) " + f"where {escape_sqlite(column)} is not null " + f"and {escape_sqlite(column)} != '' " + f"and json_array_length({escape_sqlite(column)}) > 0 " "limit 100" - ).format(column=escape_sqlite(column), sql=self.sql), + ), self.params, truncate=False, custom_time_limit=self.ds.setting( @@ -388,14 +384,14 @@ class ArrayFacet(Facet): source = source_and_config["source"] column = config.get("column") or config["simple"] # https://github.com/simonw/datasette/issues/448 - facet_sql = """ - with inner as ({sql}), + facet_sql = f""" + with inner as ({self.sql}), deduped_array_items as ( select distinct j.value, inner.* from - json_each([inner].{col}) j + json_each([inner].{escape_sqlite(column)}) j join inner ) select @@ -406,12 +402,8 @@ class ArrayFacet(Facet): group by value order by - count(*) desc, value limit {limit} - """.format( - col=escape_sqlite(column), - sql=self.sql, - limit=facet_size + 1, - ) + count(*) desc, value limit {facet_size + 1} + """ try: facet_rows_results = await self.ds.execute( self.database, diff --git a/datasette/filters.py b/datasette/filters.py index 95cc5f37..1d4e32c2 100644 --- a/datasette/filters.py +++ b/datasette/filters.py @@ -1,8 +1,11 @@ +import json +from typing import ClassVar + from datasette import hookimpl from datasette.resources import DatabaseResource -from datasette.views.base import DatasetteError from datasette.utils.asgi import BadRequest -import json +from datasette.views.base import DatasetteError + from .utils import detect_json1, escape_sqlite, path_with_removed_args @@ -99,9 +102,9 @@ def search_filters(request, database, table, datasette): fts_table=escape_sqlite(fts_table), search_col=escape_sqlite(search_col), match_clause=( - ":search_{}".format(i) + f":search_{i}" if search_mode_raw - else "escape_fts(:search_{})".format(i) + else f"escape_fts(:search_{i})" ), ) ) @@ -134,11 +137,11 @@ def through_filters(request, database, table, datasette): value = through_data["value"] db = datasette.get_database(database) outgoing_foreign_keys = await db.foreign_keys_for_table(through_table) - try: - fk_to_us = [ - fk for fk in outgoing_foreign_keys if fk["other_table"] == table - ][0] - except IndexError: + fk_to_us = next( + (fk for fk in outgoing_foreign_keys if fk["other_table"] == table), + None, + ) + if fk_to_us is None: raise DatasetteError( "Invalid _through - could not find corresponding foreign key" ) @@ -365,7 +368,7 @@ class Filters: ), ] ) - _filters_by_key = {f.key: f for f in _filters} + _filters_by_key: ClassVar[dict[str, Filter]] = {f.key: f for f in _filters} def __init__(self, pairs): self.pairs = pairs diff --git a/datasette/fixtures.py b/datasette/fixtures.py index 7c85e16a..049e35ed 100644 --- a/datasette/fixtures.py +++ b/datasette/fixtures.py @@ -1,9 +1,10 @@ -from datasette.utils.sqlite import sqlite3 -from datasette.utils import documented import itertools import random import string +from datasette.utils import documented +from datasette.utils.sqlite import sqlite3 + __all__ = [ "EXTRA_DATABASE_SQL", "TABLES", @@ -346,9 +347,7 @@ CREATE VIEW searchable_view_configured_by_metadata AS + '\nINSERT INTO no_primary_key VALUES ("RENDER_CELL_DEMO", "a202", "b202", "c202");\n' + "\n".join( [ - 'INSERT INTO compound_three_primary_keys VALUES ("{a}", "{b}", "{c}", "{content}");'.format( - a=a, b=b, c=c, content=content - ) + f'INSERT INTO compound_three_primary_keys VALUES ("{a}", "{b}", "{c}", "{content}");' for a, b, c, content in generate_compound_rows(1001) ] ) diff --git a/datasette/forbidden.py b/datasette/forbidden.py index 91b1ff96..67bf0d8b 100644 --- a/datasette/forbidden.py +++ b/datasette/forbidden.py @@ -1,4 +1,5 @@ -from datasette import hookimpl, Response +from datasette import Response, hookimpl + from .utils import add_cors_headers diff --git a/datasette/handle_exception.py b/datasette/handle_exception.py index e255ddf2..c36d5dbe 100644 --- a/datasette/handle_exception.py +++ b/datasette/handle_exception.py @@ -1,16 +1,21 @@ -from datasette import hookimpl, Response +import traceback + +from markupsafe import Markup + +from datasette import Response, hookimpl + from .utils import add_cors_headers, error_body from .utils.asgi import ( Base400, ) from .views.base import DatasetteError -from markupsafe import Markup -import traceback +# Debugger imports are deliberate - they back the "pdb" setting, which drops +# into a debugger on unhandled exceptions try: - import ipdb as pdb + import ipdb as pdb # noqa: T100 except ImportError: - import pdb + import pdb # noqa: T100 try: import rich @@ -69,7 +74,7 @@ def handle_exception(datasette, request, exception): dict( info, urls=datasette.urls, - menu_links=lambda: [], + menu_links=list, ) ), status=status, diff --git a/datasette/hookspecs.py b/datasette/hookspecs.py index 7c56f882..f89f2f36 100644 --- a/datasette/hookspecs.py +++ b/datasette/hookspecs.py @@ -1,5 +1,4 @@ -from pluggy import HookimplMarker -from pluggy import HookspecMarker +from pluggy import HookimplMarker, HookspecMarker hookspec = HookspecMarker("datasette") hookimpl = HookimplMarker("datasette") diff --git a/datasette/inspect.py b/datasette/inspect.py index 5e681e03..b126ce5c 100644 --- a/datasette/inspect.py +++ b/datasette/inspect.py @@ -1,13 +1,13 @@ import hashlib from .utils import ( - detect_spatialite, detect_fts, detect_primary_keys, + detect_spatialite, escape_sqlite, get_all_foreign_keys, - table_columns, sqlite3, + table_columns, ) HASH_BLOCK_SIZE = 1024 * 1024 @@ -95,10 +95,10 @@ def inspect_tables(conn, database_metadata): """) ] - for t in tables.keys(): + for t, table_info in tables.items(): for hidden_table in hidden_tables: if t == hidden_table or t.startswith(hidden_table): - tables[t]["hidden"] = True + table_info["hidden"] = True continue return tables diff --git a/datasette/jump.py b/datasette/jump.py index d138e827..d70d33df 100644 --- a/datasette/jump.py +++ b/datasette/jump.py @@ -21,7 +21,7 @@ class JumpSQL: search_text: str | None = None, display_name: str | None = None, item_type: str = "menu", - ) -> "JumpSQL": + ) -> JumpSQL: if search_text is None: search_text = " ".join( text for text in (label, display_name, description) if text is not None diff --git a/datasette/permissions.py b/datasette/permissions.py index 786dc026..e03b065c 100644 --- a/datasette/permissions.py +++ b/datasette/permissions.py @@ -1,7 +1,7 @@ +import contextvars from abc import ABC, abstractmethod from dataclasses import dataclass from typing import Any, NamedTuple -import contextvars # Context variable to track when permission checks should be skipped _skip_permission_checks = contextvars.ContextVar( @@ -72,8 +72,8 @@ class Resource(ABC): ) def __repr__(self) -> str: - return "{}(parent={!r}, child={!r})".format( - self.__class__.__name__, self.parent, self.child + return ( + f"{self.__class__.__name__}(parent={self.parent!r}, child={self.child!r})" ) @property @@ -129,7 +129,6 @@ class Resource(ABC): Must return two columns: parent, child """ - pass class AllowedResource(NamedTuple): diff --git a/datasette/plugins.py b/datasette/plugins.py index ae2cb17d..9cf94079 100644 --- a/datasette/plugins.py +++ b/datasette/plugins.py @@ -1,20 +1,14 @@ import importlib +import importlib.metadata as importlib_metadata +import importlib.resources as importlib_resources import os -import pluggy -from pprint import pprint import sys +from pprint import pprint + +import pluggy + from . import hookspecs -if sys.version_info >= (3, 9): - import importlib.resources as importlib_resources -else: - import importlib_resources -if sys.version_info >= (3, 10): - import importlib.metadata as importlib_metadata -else: - import importlib_metadata - - DEFAULT_PLUGINS = ( "datasette.publish.heroku", "datasette.publish.cloudrun", @@ -85,7 +79,7 @@ if DATASETTE_LOAD_PLUGINS is not None: # Ensure name can be found in plugin_to_distinfo later: pm._plugin_distinfo.append((mod, distribution)) except importlib_metadata.PackageNotFoundError: - sys.stderr.write("Plugin {} could not be found\n".format(package_name)) + sys.stderr.write(f"Plugin {package_name} could not be found\n") # Load default plugins diff --git a/datasette/publish/cloudrun.py b/datasette/publish/cloudrun.py index 63d22fe8..9ace865b 100644 --- a/datasette/publish/cloudrun.py +++ b/datasette/publish/cloudrun.py @@ -1,15 +1,17 @@ -from datasette import hookimpl -import click import json import os import re from subprocess import CalledProcessError, check_call, check_output +import click + +from datasette import hookimpl + +from ..utils import temporary_docker_directory from .common import ( add_common_publish_arguments_and_options, fail_if_publish_binary_not_installed, ) -from ..utils import temporary_docker_directory @hookimpl @@ -219,7 +221,7 @@ def publish_subcommand(publish): check_call( "gcloud builds submit --tag {}{}".format( - image_id, " --timeout {}".format(timeout) if timeout else "" + image_id, f" --timeout {timeout}" if timeout else "" ), shell=True, ) @@ -231,7 +233,7 @@ def publish_subcommand(publish): ("--min-instances", min_instances), ): if value is not None: - extra_deploy_options.append("{} {}".format(option, value)) + extra_deploy_options.append(f"{option} {value}") check_call( "gcloud run deploy --allow-unauthenticated --platform=managed --image {} {}{}".format( image_id, @@ -258,24 +260,16 @@ def _ensure_artifact_registry(artifact_project, artifact_region, artifact_reposi ) from exc describe_cmd = ( - "gcloud artifacts repositories describe {repo} --project {project} " - "--location {location} --quiet" - ).format( - repo=artifact_repository, - project=artifact_project, - location=artifact_region, + f"gcloud artifacts repositories describe {artifact_repository} --project {artifact_project} " + f"--location {artifact_region} --quiet" ) try: check_call(describe_cmd, shell=True) return except CalledProcessError: create_cmd = ( - "gcloud artifacts repositories create {repo} --repository-format=docker " - '--location {location} --project {project} --description "Datasette Cloud Run images" --quiet' - ).format( - repo=artifact_repository, - location=artifact_region, - project=artifact_project, + f"gcloud artifacts repositories create {artifact_repository} --repository-format=docker " + f'--location {artifact_region} --project {artifact_project} --description "Datasette Cloud Run images" --quiet' ) try: check_call(create_cmd, shell=True) diff --git a/datasette/publish/common.py b/datasette/publish/common.py index 29665eb3..27dfd4bf 100644 --- a/datasette/publish/common.py +++ b/datasette/publish/common.py @@ -1,9 +1,11 @@ -from ..utils import StaticMount -import click import os import shutil import sys +import click + +from ..utils import StaticMount + def add_common_publish_arguments_and_options(subcommand): for decorator in reversed( @@ -76,9 +78,7 @@ def fail_if_publish_binary_not_installed(binary, publish_target, install_link): """Exit (with error message) if ``binary` isn't installed""" if not shutil.which(binary): click.secho( - "Publishing to {publish_target} requires {binary} to be installed and configured".format( - publish_target=publish_target, binary=binary - ), + f"Publishing to {publish_target} requires {binary} to be installed and configured", bg="red", fg="white", bold=True, diff --git a/datasette/publish/heroku.py b/datasette/publish/heroku.py index f576a346..b0290833 100644 --- a/datasette/publish/heroku.py +++ b/datasette/publish/heroku.py @@ -1,19 +1,21 @@ -from contextlib import contextmanager -from datasette import hookimpl -import click import json import os import pathlib import shlex import shutil -from subprocess import call, check_output import tempfile +from contextlib import contextmanager +from subprocess import call, check_output + +import click + +from datasette import hookimpl +from datasette.utils import link_or_copy, link_or_copy_directory, parse_metadata from .common import ( add_common_publish_arguments_and_options, fail_if_publish_binary_not_installed, ) -from datasette.utils import link_or_copy, link_or_copy_directory, parse_metadata @hookimpl @@ -234,7 +236,7 @@ def temporary_heroku_directory( extras.extend(["--static", f"{mount_point}:{mount_point}"]) quoted_files = " ".join( - ["-i {}".format(shlex.quote(file_name)) for file_name in file_names] + [f"-i {shlex.quote(file_name)}" for file_name in file_names] ) procfile_cmd = "web: datasette serve --host 0.0.0.0 {quoted_files} --cors --port $PORT --inspect-file inspect-data.json {extras}".format( quoted_files=quoted_files, extras=" ".join(extras) diff --git a/datasette/renderer.py b/datasette/renderer.py index 7c94f6ee..0e01f52f 100644 --- a/datasette/renderer.py +++ b/datasette/renderer.py @@ -1,12 +1,13 @@ import json + from datasette.extras import extra_names_from_request from datasette.utils import ( - error_body, - value_as_boolean, - remove_infinites, CustomJSONEncoder, + error_body, path_from_row_pks, + remove_infinites, sqlite3, + value_as_boolean, ) from datasette.utils.asgi import Response diff --git a/datasette/static/app.css b/datasette/static/app.css index 79fb1a73..d101e4b7 100644 --- a/datasette/static/app.css +++ b/datasette/static/app.css @@ -497,10 +497,7 @@ table.rows-and-columns td em { color: #aaa; } table.rows-and-columns th { - background: #F8FAFB; padding-right: 1em; - position: relative; - z-index: 1; } table.rows-and-columns a:link { text-decoration: none; diff --git a/datasette/static/table.js b/datasette/static/table.js index 2bc9fda6..74a96d8e 100644 --- a/datasette/static/table.js +++ b/datasette/static/table.js @@ -506,45 +506,6 @@ function renderActionLink(itemConfig) { return newLink; } -function initializeStickyTableHeader(tableWrapper) { - var table = tableWrapper.querySelector("table.rows-and-columns"); - var tableHead = table && table.querySelector("thead"); - if (!tableHead || !window.requestAnimationFrame) { - return; - } - - var animationFrame = null; - function updatePosition() { - animationFrame = null; - if (window.matchMedia("(max-width: 576px)").matches) { - tableHead.style.transform = ""; - return; - } - // overflow-x: auto makes tableWrapper the sticky scroll container, even - // though the page handles vertical scrolling. Counter the page scroll so - // the real header stays aligned with both the viewport and the table. - var wrapperTop = tableWrapper.getBoundingClientRect().top; - var maximumOffset = Math.max( - 0, - table.offsetHeight - tableHead.offsetHeight, - ); - var offset = Math.min(Math.max(-wrapperTop, 0), maximumOffset); - tableHead.style.transform = offset - ? `translateY(${offset}px)` - : ""; - } - - function queuePositionUpdate() { - if (animationFrame === null) { - animationFrame = window.requestAnimationFrame(updatePosition); - } - } - - window.addEventListener("scroll", queuePositionUpdate, { passive: true }); - window.addEventListener("resize", queuePositionUpdate); - queuePositionUpdate(); -} - /** Main initialization function for Datasette Table interactions */ const initDatasetteTable = function (manager) { // Feature detection @@ -560,7 +521,6 @@ const initDatasetteTable = function (manager) { if (tableWrapper) { tableWrapper.addEventListener("scroll", closeMenu); } - window.addEventListener("scroll", closeMenu, { passive: true }); document.body.addEventListener("click", (ev) => { /* was this click outside the menu? */ var target = ev.target; @@ -671,9 +631,6 @@ const initDatasetteTable = function (manager) { icon.addEventListener("click", onTableHeaderClick); th.appendChild(icon); }); - if (tableWrapper) { - initializeStickyTableHeader(tableWrapper); - } }; function filterRowSelector(manager) { diff --git a/datasette/stored_queries.py b/datasette/stored_queries.py index f5f977d9..db3c6548 100644 --- a/datasette/stored_queries.py +++ b/datasette/stored_queries.py @@ -1,8 +1,9 @@ from __future__ import annotations -from dataclasses import dataclass import json -from typing import Any, Iterable +from collections.abc import Iterable +from dataclasses import dataclass +from typing import Any from .utils import tilde_encode, urlsafe_components @@ -386,7 +387,7 @@ async def count_queries( OR q.sql LIKE :query_search ) """) - params["query_search"] = "%{}%".format(q) + params["query_search"] = f"%{q}%" if is_write is not None: where_clauses.append("q.is_write = :query_is_write") params["query_is_write"] = int(bool(is_write)) @@ -462,7 +463,7 @@ async def list_queries( except ValueError: components = [] if database is None and len(components) == 3: - where_clauses.append(""" + where_clauses.append(f""" ( q.database_name > :cursor_database OR ( @@ -476,12 +477,12 @@ async def list_queries( ) ) ) - """.format(sort_key_sql=sort_key_sql)) + """) params["cursor_database"] = components[0] params["cursor_sort_key"] = components[1] params["cursor_name"] = components[2] elif database is not None and len(components) == 2: - where_clauses.append(""" + where_clauses.append(f""" ( {sort_key_sql} > :cursor_sort_key OR ( @@ -489,7 +490,7 @@ async def list_queries( AND q.name > :cursor_name ) ) - """.format(sort_key_sql=sort_key_sql)) + """) params["cursor_sort_key"] = components[0] params["cursor_name"] = components[1] @@ -502,7 +503,7 @@ async def list_queries( OR q.sql LIKE :query_search ) """) - params["query_search"] = "%{}%".format(q) + params["query_search"] = f"%{q}%" if is_write is not None: where_clauses.append("q.is_write = :query_is_write") params["query_is_write"] = int(bool(is_write)) diff --git a/datasette/templates/_permission_ui_styles.html b/datasette/templates/_permission_ui_styles.html index 53a824f1..21a2ea8f 100644 --- a/datasette/templates/_permission_ui_styles.html +++ b/datasette/templates/_permission_ui_styles.html @@ -6,8 +6,20 @@ padding: 1.5em; margin-bottom: 2em; } +.permission-form form { + max-width: 60rem; +} +.permission-form-grid { + display: grid; + gap: 1.5rem; + grid-template-columns: repeat(2, minmax(0, 1fr)); +} +.permission-form-result { + margin-top: 1rem; + max-width: 60rem; +} .form-section { - margin-bottom: 1em; + margin-bottom: 1.25em; } .form-section label { display: block; @@ -15,22 +27,51 @@ font-weight: bold; } .form-section input[type="text"], -.form-section select { - width: 100%; - max-width: 500px; - padding: 0.5em; +.form-section input[type="number"], +.form-section select, +.permission-textarea { + background-color: #fff; + border: 1px solid #aaa; + border-radius: 4px; box-sizing: border-box; - border: 1px solid #ccc; - border-radius: 3px; + box-shadow: inset 0 1px 2px rgba(0, 0, 0, 0.08); + color: #222; + font-family: inherit; + font-size: 1rem; + line-height: 1.4; + max-width: none; + width: 100%; +} +.form-section input[type="text"] { + height: 3rem; + padding: 0.6rem 0.75rem; +} +.form-section input[type="number"] { + height: 3rem; + max-width: 7rem; + padding: 0.6rem 0.75rem; +} +.form-section select { + height: 3rem; + padding: 0.6rem 0.75rem; +} +.permission-textarea { + font-family: monospace; + min-height: 12rem; + padding: 0.75rem; + resize: vertical; } .form-section input[type="text"]:focus, -.form-section select:focus { - outline: 2px solid #0066cc; +.form-section input[type="number"]:focus, +.form-section select:focus, +.permission-textarea:focus { border-color: #0066cc; + box-shadow: 0 0 0 3px rgba(0, 102, 204, 0.18); + outline: none; } .form-section small { display: block; - margin-top: 0.3em; + margin-top: 0.45em; color: #666; } .form-actions { @@ -142,4 +183,9 @@ text-align: center; color: #666; } +@media only screen and (max-width: 576px) { + .permission-form-grid { + grid-template-columns: minmax(0, 1fr); + } +} diff --git a/datasette/templates/_permissions_debug_tabs.html b/datasette/templates/_permissions_debug_tabs.html index d7203c1e..8e0f486e 100644 --- a/datasette/templates/_permissions_debug_tabs.html +++ b/datasette/templates/_permissions_debug_tabs.html @@ -44,10 +44,10 @@ diff --git a/datasette/templates/allow_debug.html b/datasette/templates/allow_debug.html index 1ecc92df..fda4032c 100644 --- a/datasette/templates/allow_debug.html +++ b/datasette/templates/allow_debug.html @@ -3,29 +3,11 @@ {% block title %}Debug allow rules{% endblock %} {% block extra_head %} +{% include "_permission_ui_styles.html" %} {% endblock %} @@ -38,24 +20,28 @@ p.message-warning {

Use this tool to try out different actor and allow combinations. See Defining permissions with "allow" blocks for documentation.

-
-
-

- -
-
-

- -
-
- -
-
+
+
+
+
+ + +
+
+ + +
+
+
+ +
+
-{% if error %}

{{ error }}

{% endif %} + {% if error %}

{{ error }}

{% endif %} -{% if result == "True" %}

Result: allow

{% endif %} + {% if result == "True" %}

Result: allow

{% endif %} -{% if result == "False" %}

Result: deny

{% endif %} + {% if result == "False" %}

Result: deny

{% endif %} +
{% endblock %} diff --git a/datasette/templates/debug_allowed.html b/datasette/templates/debug_allowed.html index 83cc1ae6..80249d9c 100644 --- a/datasette/templates/debug_allowed.html +++ b/datasette/templates/debug_allowed.html @@ -49,7 +49,7 @@
- + Number of results per page (max 200)
diff --git a/datasette/templates/debug_check.html b/datasette/templates/debug_check.html index 3b229a25..b9fc636a 100644 --- a/datasette/templates/debug_check.html +++ b/datasette/templates/debug_check.html @@ -1,6 +1,6 @@ {% extends "base.html" %} -{% block title %}Permission Check{% endblock %} +{% block title %}Explain a permission decision{% endblock %} {% block extra_head %} @@ -13,29 +13,35 @@ border-radius: 5px; } #output.allowed { - background-color: #e8f5e9; + background-color: #f3fbf4; border: 2px solid #4caf50; } #output.denied { - background-color: #ffebee; + background-color: #fff7f7; border: 2px solid #f44336; } #output h2 { margin-top: 0; } -#output .result-badge { +#output h3 { + margin-bottom: 0.5em; +} +#output .result-badge, +.effect-badge, +.rule-status { display: inline-block; - padding: 0.3em 0.8em; + padding: 0.2em 0.5em; border-radius: 3px; font-weight: bold; - font-size: 1.1em; } -#output .allowed-badge { - background-color: #4caf50; +#output .allowed-badge, +.effect-allow { + background-color: #2e7d32; color: white; } -#output .denied-badge { - background-color: #f44336; +#output .denied-badge, +.effect-deny { + background-color: #c62828; color: white; } .details-section { @@ -48,70 +54,130 @@ .details-section dd { margin-left: 1em; } +.explanation-section { + background: rgba(255, 255, 255, 0.75); + border: 1px solid #ddd; + border-radius: 4px; + margin-top: 1em; + padding: 0 1em 1em; +} +.rules-table { + border-collapse: collapse; + width: 100%; +} +.rules-table th, +.rules-table td { + border-bottom: 1px solid #ddd; + padding: 0.5em; + text-align: left; + vertical-align: top; +} +.rule-status { + background: #e8f5e9; + color: #1b5e20; +} +.rule-ignored { + background: #eee; + color: #555; + font-weight: normal; +} +.requirement-allowed { + color: #1b5e20; +} +.requirement-denied { + color: #b71c1c; +} +@media only screen and (max-width: 576px) { + .rules-table, + .rules-table tbody, + .rules-table tr, + .rules-table td { + display: block; + } + .rules-table thead { + display: none; + } + .rules-table td::before { + content: attr(data-label) ": "; + font-weight: bold; + } +} {% endblock %} {% block content %} -

Permission check

+

Explain a permission decision

{% set current_tab = "check" %} {% include "_permissions_debug_tabs.html" %} -

Use this tool to test permission checks for the current actor. It queries the /-/check.json API endpoint.

- -{% if request.actor %} -

Current actor: {{ request.actor.get("id", "anonymous") }}

-{% else %} -

Current actor: anonymous (not logged in)

-{% endif %} +

Test an actor, action and resource. The result explains which rules matched, which specificity level won, and whether actor restrictions or required actions changed the verdict.

-
+
- + + + Use null for an anonymous actor. This actor is simulated; it does not change who you are signed in as. +
+ +
+ - The permission action to check + The operation to evaluate
-
- +
+ - For database-level permissions, specify the database name + The database or other parent resource
-
- - - For table-level permissions, specify the table name (requires parent) +
+ + + The table, query or other child resource
- +
+actionSelect.addEventListener('change', updateResourceFields); +(function initializeFromUrl() { + const params = populateFormFromURL(); + updateResourceFields(); + if (params.get('action')) { + performCheck(); + } +})(); + {% endblock %} diff --git a/datasette/templates/debug_permissions_playground.html b/datasette/templates/debug_permissions_playground.html index 4410a677..8b0cbbcf 100644 --- a/datasette/templates/debug_permissions_playground.html +++ b/datasette/templates/debug_permissions_playground.html @@ -1,6 +1,6 @@ {% extends "base.html" %} -{% block title %}Debug permissions{% endblock %} +{% block title %}Permission activity{% endblock %} {% block extra_head %} {% include "_permission_ui_styles.html" %} @@ -20,60 +20,45 @@ .check-action, .check-when, .check-result { font-size: 1.3em; } -textarea { - height: 10em; - width: 95%; - box-sizing: border-box; - padding: 0.5em; - border: 2px dotted black; -} -.two-col { - display: inline-block; - width: 48%; -} -.two-col label { - width: 48%; -} -@media only screen and (max-width: 576px) { - .two-col { - width: 100%; - } -} {% endblock %} {% block content %} -

Permission playground

+

Permission activity

{% set current_tab = "permissions" %} {% include "_permissions_debug_tabs.html" %} -

This tool lets you simulate an actor and a permission check for that actor.

+

Raw simulator

+ +

This form runs a hypothetical permission check and returns its raw explanation JSON. Use the Explain tool for a visual explanation of the same decision.

-
-
- - +
+
+
+ + +
-
-
-
- - -
-
- - -
-
- - +
+
+ + +
+
+ + +
+
+ + +
@@ -125,7 +110,7 @@ debugPost.addEventListener('submit', function(ev) { }); -

Recent permissions checks

+

Recent permission checks

{% if filter != "all" %}All{% else %}All{% endif %}, diff --git a/datasette/templates/debug_rules.html b/datasette/templates/debug_rules.html index d00ba9cc..233c0e94 100644 --- a/datasette/templates/debug_rules.html +++ b/datasette/templates/debug_rules.html @@ -37,7 +37,7 @@

- + Number of results per page (max 200)
diff --git a/datasette/tokens.py b/datasette/tokens.py index 4f905339..79f840d2 100644 --- a/datasette/tokens.py +++ b/datasette/tokens.py @@ -10,7 +10,7 @@ from __future__ import annotations import dataclasses import time -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING import itsdangerous @@ -50,24 +50,24 @@ class TokenRestrictions: database: dict[str, list[str]] = dataclasses.field(default_factory=dict) resource: dict[str, dict[str, list[str]]] = dataclasses.field(default_factory=dict) - def allow_all(self, action: str) -> "TokenRestrictions": + def allow_all(self, action: str) -> TokenRestrictions: """Allow an action across all databases and resources.""" self.all.append(action) return self - def allow_database(self, database: str, action: str) -> "TokenRestrictions": + def allow_database(self, database: str, action: str) -> TokenRestrictions: """Allow an action on a specific database.""" self.database.setdefault(database, []).append(action) return self def allow_resource( self, database: str, resource: str, action: str - ) -> "TokenRestrictions": + ) -> TokenRestrictions: """Allow an action on a specific resource within a database.""" self.resource.setdefault(database, {}).setdefault(resource, []).append(action) return self - def abbreviated(self, datasette: "Datasette") -> Optional[dict]: + def abbreviated(self, datasette: Datasette) -> dict | None: """ Return the abbreviated ``_r`` dictionary shape for this set of restrictions, using action abbreviations registered with ``datasette``. @@ -112,16 +112,16 @@ class TokenHandler: async def create_token( self, - datasette: "Datasette", + datasette: Datasette, actor_id: str, *, - expires_after: Optional[int] = None, - restrictions: Optional[TokenRestrictions] = None, + expires_after: int | None = None, + restrictions: TokenRestrictions | None = None, ) -> str: """Create and return a token string for the given actor.""" raise NotImplementedError - async def verify_token(self, datasette: "Datasette", token: str) -> Optional[dict]: + async def verify_token(self, datasette: Datasette, token: str) -> dict | None: """ Verify a token and return an actor dict. @@ -142,11 +142,11 @@ class SignedTokenHandler(TokenHandler): async def create_token( self, - datasette: "Datasette", + datasette: Datasette, actor_id: str, *, - expires_after: Optional[int] = None, - restrictions: Optional[TokenRestrictions] = None, + expires_after: int | None = None, + restrictions: TokenRestrictions | None = None, ) -> str: if not datasette.setting("allow_signed_tokens"): raise ValueError( @@ -163,7 +163,7 @@ class SignedTokenHandler(TokenHandler): token["_r"] = abbreviated return "dstok_{}".format(datasette.sign(token, namespace="token")) - async def verify_token(self, datasette: "Datasette", token: str) -> Optional[dict]: + async def verify_token(self, datasette: Datasette, token: str) -> dict | None: prefix = "dstok_" if not token.startswith(prefix): @@ -200,9 +200,8 @@ class SignedTokenHandler(TokenHandler): ): duration = max_signed_tokens_ttl - if duration: - if time.time() - created > duration: - raise TokenInvalid("Token has expired") + if duration and time.time() - created > duration: + raise TokenInvalid("Token has expired") actor = {"id": decoded["a"], "token": "dstok"} diff --git a/datasette/tracer.py b/datasette/tracer.py index 28f3cc09..1fbda6f9 100644 --- a/datasette/tracer.py +++ b/datasette/tracer.py @@ -1,10 +1,11 @@ import asyncio +import json +import time +import traceback from contextlib import contextmanager from contextvars import ContextVar + from markupsafe import escape -import time -import json -import traceback tracers = {} @@ -132,17 +133,17 @@ class AsgiTracer: "num_traces": len(traces), "traces": traces, } - try: - content_type = [ + content_type = next( + ( v.decode("utf8") for k, v in response_headers if k.lower() == b"content-type" - ][0] - except IndexError: - content_type = "" + ), + "", + ) if "text/html" in content_type and b"" in accumulated_body: extra = escape(json.dumps(trace_info, indent=2)) - extra_html = f"
{extra}
".encode("utf8") + extra_html = f"
{extra}
".encode() accumulated_body = accumulated_body.replace(b"", extra_html) elif "json" in content_type and accumulated_body.startswith(b"{"): data = json.loads(accumulated_body.decode("utf8")) diff --git a/datasette/url_builder.py b/datasette/url_builder.py index 16b3d42b..f8da20f3 100644 --- a/datasette/url_builder.py +++ b/datasette/url_builder.py @@ -1,6 +1,7 @@ -from .utils import tilde_encode, path_with_format, PrefixedUrlString import urllib +from .utils import PrefixedUrlString, path_with_format, tilde_encode + class Urls: def __init__(self, ds): @@ -8,8 +9,7 @@ class Urls: def path(self, path, format=None): if not isinstance(path, PrefixedUrlString): - if path.startswith("/"): - path = path[1:] + path = path.removeprefix("/") path = self.ds.setting("base_url") + path if format is not None: path = path_with_format(path=path, format=format) @@ -56,6 +56,7 @@ class Urls: return PrefixedUrlString(path) def row_blob(self, database, table, row_path, column): - return self.table(database, table) + "/{}.blob?_blob_column={}".format( - row_path, urllib.parse.quote_plus(column) + return ( + self.table(database, table) + + f"/{row_path}.blob?_blob_column={urllib.parse.quote_plus(column)}" ) diff --git a/datasette/utils/__init__.py b/datasette/utils/__init__.py index 42574d3b..eb46d549 100644 --- a/datasette/utils/__init__.py +++ b/datasette/utils/__init__.py @@ -1,29 +1,31 @@ import asyncio +import base64 import binascii -from contextlib import contextmanager -import aiofiles -import click -from collections import OrderedDict, namedtuple, Counter import copy import dataclasses -import base64 import hashlib import inspect import json -import markupsafe -import mergedeep import os import re +import secrets import shlex +import shutil import tempfile -import typing import time import types -import secrets -import shutil -from typing import Iterable, List, Tuple +import typing import urllib +from collections import Counter, OrderedDict, namedtuple +from collections.abc import Iterable +from contextlib import contextmanager + +import aiofiles +import click +import markupsafe +import mergedeep import yaml + from .shutil_backport import copytree from .sqlite import sqlite3, supports_table_xinfo @@ -36,7 +38,7 @@ if typing.TYPE_CHECKING: class PaginatedResources: """Paginated results from allowed_resources query.""" - resources: List["Resource"] + resources: list["Resource"] next: str | None # Keyset token for next page (None if no more results) _datasette: typing.Any = dataclasses.field(default=None, repr=False) _action: str = dataclasses.field(default=None, repr=False) @@ -83,22 +85,132 @@ class PaginatedResources: # From https://www.sqlite.org/lang_keywords.html -reserved_words = set( - ( - "abort action add after all alter analyze and as asc attach autoincrement " - "before begin between by cascade case cast check collate column commit " - "conflict constraint create cross current_date current_time " - "current_timestamp database default deferrable deferred delete desc detach " - "distinct drop each else end escape except exclusive exists explain fail " - "for foreign from full glob group having if ignore immediate in index " - "indexed initially inner insert instead intersect into is isnull join key " - "left like limit match natural no not notnull null of offset on or order " - "outer plan pragma primary query raise recursive references regexp reindex " - "release rename replace restrict right rollback row savepoint select set " - "table temp temporary then to transaction trigger union unique update using " - "vacuum values view virtual when where with without" - ).split() -) +reserved_words = { + "abort", + "action", + "add", + "after", + "all", + "alter", + "analyze", + "and", + "as", + "asc", + "attach", + "autoincrement", + "before", + "begin", + "between", + "by", + "cascade", + "case", + "cast", + "check", + "collate", + "column", + "commit", + "conflict", + "constraint", + "create", + "cross", + "current_date", + "current_time", + "current_timestamp", + "database", + "default", + "deferrable", + "deferred", + "delete", + "desc", + "detach", + "distinct", + "drop", + "each", + "else", + "end", + "escape", + "except", + "exclusive", + "exists", + "explain", + "fail", + "for", + "foreign", + "from", + "full", + "glob", + "group", + "having", + "if", + "ignore", + "immediate", + "in", + "index", + "indexed", + "initially", + "inner", + "insert", + "instead", + "intersect", + "into", + "is", + "isnull", + "join", + "key", + "left", + "like", + "limit", + "match", + "natural", + "no", + "not", + "notnull", + "null", + "of", + "offset", + "on", + "or", + "order", + "outer", + "plan", + "pragma", + "primary", + "query", + "raise", + "recursive", + "references", + "regexp", + "reindex", + "release", + "rename", + "replace", + "restrict", + "right", + "rollback", + "row", + "savepoint", + "select", + "set", + "table", + "temp", + "temporary", + "then", + "to", + "transaction", + "trigger", + "union", + "unique", + "update", + "using", + "vacuum", + "values", + "view", + "virtual", + "when", + "where", + "with", + "without", +} APT_GET_DOCKERFILE_EXTRAS = r""" RUN apt-get update && \ @@ -158,7 +270,7 @@ functions_marked_as_documented = [] def documented(fn=None, *, label=None): def decorate(fn): - fn._datasette_docs_label = label or "internals_utils_{}".format(fn.__name__) + fn._datasette_docs_label = label or f"internals_utils_{fn.__name__}" functions_marked_as_documented.append(fn) return fn @@ -360,7 +472,7 @@ disallawed_sql_res = [ ( re.compile(f"pragma(?!_({'|'.join(allowed_pragmas)}))"), "Statement contained a disallowed PRAGMA. Allowed pragma functions are {}".format( - ", ".join("pragma_{}()".format(pragma) for pragma in allowed_pragmas) + ", ".join(f"pragma_{pragma}()" for pragma in allowed_pragmas) ), ) ] @@ -534,10 +646,7 @@ CMD {cmd}""".format( else "" ), environment_variables="\n".join( - [ - "ENV {} '{}'".format(key, value) - for key, value in environment_variables.items() - ] + [f"ENV {key} '{value}'" for key, value in environment_variables.items()] ), install_from=" ".join(install), files=" ".join(files), @@ -636,11 +745,11 @@ def detect_primary_keys(conn, table): def get_outbound_foreign_keys(conn, table): - infos = conn.execute(f"PRAGMA foreign_key_list([{table}])").fetchall() + infos = conn.execute(f"PRAGMA foreign_key_list({escape_sqlite(table)})").fetchall() fks = [] for info in infos: if info is not None: - id, seq, table_name, from_, to_, on_update, on_delete, match = info + id, seq, table_name, from_, to_, _on_update, _on_delete, _match = info fks.append( { "column": from_, @@ -741,7 +850,7 @@ def detect_json1(conn=None): try: conn.execute("SELECT json('{}')") return True - except Exception: + except sqlite3.Error: return False finally: if close_conn: @@ -821,9 +930,7 @@ def is_url(value): if not value.startswith("http://") and not value.startswith("https://"): return False # Any whitespace at all is invalid - if whitespace_re.search(value): - return False - return True + return not whitespace_re.search(value) css_class_re = re.compile(r"^[a-zA-Z]+[_a-zA-Z0-9-]*$") @@ -876,7 +983,9 @@ def module_from_path(path, name): mod.__file__ = path with open(path, "r") as file: code = compile(file.read(), path, "exec", dont_inherit=True) - exec(code, mod.__dict__) + # Executing the file is the whole point - this is how --plugins-dir loads + # plugins and how metadata/config .py files are evaluated + exec(code, mod.__dict__) # noqa: S102 return mod @@ -1033,9 +1142,7 @@ def escape_fts(query): query += '"' bits = _escape_fts_re.split(query) bits = [b for b in bits if b and b != '""'] - return " ".join( - '"{}"'.format(bit) if not bit.startswith('"') else bit for bit in bits - ) + return " ".join(f'"{bit}"' if not bit.startswith('"') else bit for bit in bits) class MultiParams: @@ -1047,7 +1154,7 @@ class MultiParams: data[key], (list, tuple) ), "dictionary data should be a dictionary of key => [list]" self._data = data - elif isinstance(data, list) or isinstance(data, tuple): + elif isinstance(data, (list, tuple)): new_data = {} for item in data: assert ( @@ -1137,9 +1244,7 @@ def _gather_arguments(fn, kwargs): for parameter in parameters: if parameter not in kwargs: raise TypeError( - "{} requires parameters {}, missing: {}".format( - fn, tuple(parameters), set(parameters) - set(kwargs.keys()) - ) + f"{fn} requires parameters {tuple(parameters)}, missing: {set(parameters) - set(kwargs.keys())}" ) call_with.append(kwargs[parameter]) return call_with @@ -1208,9 +1313,9 @@ def resolve_env_secrets(config, environ): """Create copy that recursively replaces {"$env": "NAME"} with values from environ""" if isinstance(config, dict): if list(config.keys()) == ["$env"]: - return environ.get(list(config.values())[0]) + return environ.get(next(iter(config.values()))) elif list(config.keys()) == ["$file"]: - with open(list(config.values())[0]) as fp: + with open(next(iter(config.values()))) as fp: return fp.read() else: return { @@ -1306,7 +1411,7 @@ _named_param_re = re.compile(r":(\w+)") @documented -def named_parameters(sql: str) -> List[str]: +def named_parameters(sql: str) -> list[str]: """ Given a SQL statement, return a list of named parameters that are used in the statement @@ -1319,7 +1424,7 @@ def named_parameters(sql: str) -> List[str]: return _named_param_re.findall(sql) -async def derive_named_parameters(db: "Database", sql: str) -> List[str]: +async def derive_named_parameters(db: "Database", sql: str) -> list[str]: """ This undocumented but stable method exists for backwards compatibility with plugins that were using it before it switched to named_parameters() @@ -1343,9 +1448,9 @@ def parse_size_limit(value, default, maximum, name="_size"): if size < 0: raise ValueError except ValueError: - raise ValueError("{} must be a positive integer".format(name)) + raise ValueError(f"{name} must be a positive integer") if size > maximum: - raise ValueError("{} must be <= {}".format(name, maximum)) + raise ValueError(f"{name} must be <= {maximum}") return size @@ -1403,7 +1508,7 @@ class TildeEncoder(dict): elif b == _space: res = "+" else: - res = "~{:02X}".format(b) + res = f"~{b:02X}" self[b] = res return res @@ -1498,7 +1603,7 @@ def _combine(base: dict, update: dict) -> dict: return base -def pairs_to_nested_config(pairs: typing.List[typing.Tuple[str, typing.Any]]) -> dict: +def pairs_to_nested_config(pairs: list[tuple[str, typing.Any]]) -> dict: """ Parse a list of key-value pairs into a nested dictionary. """ @@ -1513,7 +1618,7 @@ def make_slot_function(name, datasette, request, **kwargs): from datasette.plugins import pm method = getattr(pm.hook, name, None) - assert method is not None, "No hook found for {}".format(name) + assert method is not None, f"No hook found for {name}" async def inner(): html_bits = [] @@ -1537,7 +1642,7 @@ def prune_empty_dicts(d: dict): d.pop(key, None) -def move_plugins_and_allow(source: dict, destination: dict) -> Tuple[dict, dict]: +def move_plugins_and_allow(source: dict, destination: dict) -> tuple[dict, dict]: """ Move 'plugins' and 'allow' keys from source to destination dictionary. Creates hierarchy in destination if needed. After moving, recursively remove any keys diff --git a/datasette/utils/actions_sql.py b/datasette/utils/actions_sql.py index c7137e6b..d767e391 100644 --- a/datasette/utils/actions_sql.py +++ b/datasette/utils/actions_sql.py @@ -252,88 +252,62 @@ async def _build_single_action_sql( ] ) - # Continue with the cascading logic - query_parts.extend( - [ - "child_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow,", - " json_group_array(CASE WHEN ar.allow = 0 THEN ar.source_plugin || ': ' || ar.reason END) AS deny_reasons,", - " json_group_array(CASE WHEN ar.allow = 1 THEN ar.source_plugin || ': ' || ar.reason END) AS allow_reasons", - " FROM base b", - " LEFT JOIN all_rules ar ON ar.parent = b.parent AND ar.child = b.child", - " GROUP BY b.parent, b.child", - "),", - "parent_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow,", - " json_group_array(CASE WHEN ar.allow = 0 THEN ar.source_plugin || ': ' || ar.reason END) AS deny_reasons,", - " json_group_array(CASE WHEN ar.allow = 1 THEN ar.source_plugin || ': ' || ar.reason END) AS allow_reasons", - " FROM base b", - " LEFT JOIN all_rules ar ON ar.parent = b.parent AND ar.child IS NULL", - " GROUP BY b.parent, b.child", - "),", - "global_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow,", - " json_group_array(CASE WHEN ar.allow = 0 THEN ar.source_plugin || ': ' || ar.reason END) AS deny_reasons,", - " json_group_array(CASE WHEN ar.allow = 1 THEN ar.source_plugin || ': ' || ar.reason END) AS allow_reasons", - " FROM base b", - " LEFT JOIN all_rules ar ON ar.parent IS NULL AND ar.child IS NULL", - " GROUP BY b.parent, b.child", - "),", + # Continue with the cascading logic. + # Aggregate the RULES by cascade level (small), rather than grouping + # base x rules (which scales with the number of resources). + def _agg(select_key, where, group_by): + parts = [ + f" SELECT {select_key}", + " MAX(CASE WHEN allow = 0 THEN 1 ELSE 0 END) AS any_deny,", + " MAX(CASE WHEN allow = 1 THEN 1 ELSE 0 END) AS any_allow,", + " json_group_array(CASE WHEN allow = 0 THEN source_plugin || ': ' || reason END) AS deny_reasons,", + " json_group_array(CASE WHEN allow = 1 THEN source_plugin || ': ' || reason END) AS allow_reasons", + f" FROM all_rules WHERE {where}", ] + if group_by: + parts.append(f" GROUP BY {group_by}") + return parts + + query_parts.extend( + ["child_agg AS ("] + + _agg( + "parent, child,", + "parent IS NOT NULL AND child IS NOT NULL", + "parent, child", + ) + + ["),", "parent_agg AS ("] + + _agg("parent,", "parent IS NOT NULL AND child IS NULL", "parent") + + ["),", "global_agg AS ("] + + _agg("", "parent IS NULL AND child IS NULL", None) + + ["),"] ) # Add anonymous decision logic if needed if include_is_private: - query_parts.extend( - [ - "anon_child_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow", - " FROM base b", - " LEFT JOIN anon_rules ar ON ar.parent = b.parent AND ar.child = b.child", - " GROUP BY b.parent, b.child", - "),", - "anon_parent_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow", - " FROM base b", - " LEFT JOIN anon_rules ar ON ar.parent = b.parent AND ar.child IS NULL", - " GROUP BY b.parent, b.child", - "),", - "anon_global_lvl AS (", - " SELECT b.parent, b.child,", - " MAX(CASE WHEN ar.allow = 0 THEN 1 ELSE 0 END) AS any_deny,", - " MAX(CASE WHEN ar.allow = 1 THEN 1 ELSE 0 END) AS any_allow", - " FROM base b", - " LEFT JOIN anon_rules ar ON ar.parent IS NULL AND ar.child IS NULL", - " GROUP BY b.parent, b.child", - "),", - "anon_decisions AS (", - " SELECT", - " b.parent, b.child,", - " CASE", - " WHEN acl.any_deny = 1 THEN 0", - " WHEN acl.any_allow = 1 THEN 1", - " WHEN apl.any_deny = 1 THEN 0", - " WHEN apl.any_allow = 1 THEN 1", - " WHEN agl.any_deny = 1 THEN 0", - " WHEN agl.any_allow = 1 THEN 1", - " ELSE 0", - " END AS anon_is_allowed", - " FROM base b", - " JOIN anon_child_lvl acl ON b.parent = acl.parent AND (b.child = acl.child OR (b.child IS NULL AND acl.child IS NULL))", - " JOIN anon_parent_lvl apl ON b.parent = apl.parent AND (b.child = apl.child OR (b.child IS NULL AND apl.child IS NULL))", - " JOIN anon_global_lvl agl ON b.parent = agl.parent AND (b.child = agl.child OR (b.child IS NULL AND agl.child IS NULL))", - "),", + + def _anon_agg(select_key, where, group_by): + parts = [ + f" SELECT {select_key}", + " MAX(CASE WHEN allow = 0 THEN 1 ELSE 0 END) AS any_deny,", + " MAX(CASE WHEN allow = 1 THEN 1 ELSE 0 END) AS any_allow", + f" FROM anon_rules WHERE {where}", ] + if group_by: + parts.append(f" GROUP BY {group_by}") + return parts + + query_parts.extend( + ["anon_child_agg AS ("] + + _anon_agg( + "parent, child,", + "parent IS NOT NULL AND child IS NOT NULL", + "parent, child", + ) + + ["),", "anon_parent_agg AS ("] + + _anon_agg("parent,", "parent IS NOT NULL AND child IS NULL", "parent") + + ["),", "anon_global_agg AS ("] + + _anon_agg("", "parent IS NULL AND child IS NULL", None) + + ["),"] ) # Final decisions @@ -342,31 +316,28 @@ async def _build_single_action_sql( "decisions AS (", " SELECT", " b.parent, b.child,", - " -- Cascading permission logic: child → parent → global, DENY beats ALLOW at each level", + " -- Cascading permission logic: child -> parent -> global, DENY beats ALLOW at each level", " -- Priority order:", - " -- 1. Child-level deny (most specific, blocks access)", - " -- 2. Child-level allow (most specific, grants access)", - " -- 3. Parent-level deny (intermediate, blocks access)", - " -- 4. Parent-level allow (intermediate, grants access)", - " -- 5. Global-level deny (least specific, blocks access)", - " -- 6. Global-level allow (least specific, grants access)", + " -- 1. Child-level deny 2. Child-level allow", + " -- 3. Parent-level deny 4. Parent-level allow", + " -- 5. Global-level deny 6. Global-level allow", " -- 7. Default deny (no rules match)", " CASE", - " WHEN cl.any_deny = 1 THEN 0", - " WHEN cl.any_allow = 1 THEN 1", - " WHEN pl.any_deny = 1 THEN 0", - " WHEN pl.any_allow = 1 THEN 1", - " WHEN gl.any_deny = 1 THEN 0", - " WHEN gl.any_allow = 1 THEN 1", + " WHEN ca.any_deny = 1 THEN 0", + " WHEN ca.any_allow = 1 THEN 1", + " WHEN pa.any_deny = 1 THEN 0", + " WHEN pa.any_allow = 1 THEN 1", + " WHEN ga.any_deny = 1 THEN 0", + " WHEN ga.any_allow = 1 THEN 1", " ELSE 0", " END AS is_allowed,", " CASE", - " WHEN cl.any_deny = 1 THEN cl.deny_reasons", - " WHEN cl.any_allow = 1 THEN cl.allow_reasons", - " WHEN pl.any_deny = 1 THEN pl.deny_reasons", - " WHEN pl.any_allow = 1 THEN pl.allow_reasons", - " WHEN gl.any_deny = 1 THEN gl.deny_reasons", - " WHEN gl.any_allow = 1 THEN gl.allow_reasons", + " WHEN ca.any_deny = 1 THEN ca.deny_reasons", + " WHEN ca.any_allow = 1 THEN ca.allow_reasons", + " WHEN pa.any_deny = 1 THEN pa.deny_reasons", + " WHEN pa.any_allow = 1 THEN pa.allow_reasons", + " WHEN ga.any_deny = 1 THEN ga.deny_reasons", + " WHEN ga.any_allow = 1 THEN ga.allow_reasons", " ELSE '[]'", " END AS reason", ] @@ -374,21 +345,34 @@ async def _build_single_action_sql( if include_is_private: query_parts.append( - " , CASE WHEN ad.anon_is_allowed = 0 THEN 1 ELSE 0 END AS is_private" + " , CASE WHEN (" + "CASE" + " WHEN aca.any_deny = 1 THEN 0" + " WHEN aca.any_allow = 1 THEN 1" + " WHEN apa.any_deny = 1 THEN 0" + " WHEN apa.any_allow = 1 THEN 1" + " WHEN aga.any_deny = 1 THEN 0" + " WHEN aga.any_allow = 1 THEN 1" + " ELSE 0 END" + ") = 0 THEN 1 ELSE 0 END AS is_private" ) query_parts.extend( [ " FROM base b", - " JOIN child_lvl cl ON b.parent = cl.parent AND (b.child = cl.child OR (b.child IS NULL AND cl.child IS NULL))", - " JOIN parent_lvl pl ON b.parent = pl.parent AND (b.child = pl.child OR (b.child IS NULL AND pl.child IS NULL))", - " JOIN global_lvl gl ON b.parent = gl.parent AND (b.child = gl.child OR (b.child IS NULL AND gl.child IS NULL))", + " LEFT JOIN child_agg ca ON ca.parent = b.parent AND ca.child = b.child", + " LEFT JOIN parent_agg pa ON pa.parent = b.parent", + " CROSS JOIN global_agg ga", ] ) if include_is_private: - query_parts.append( - " JOIN anon_decisions ad ON b.parent = ad.parent AND (b.child = ad.child OR (b.child IS NULL AND ad.child IS NULL))" + query_parts.extend( + [ + " LEFT JOIN anon_child_agg aca ON aca.parent = b.parent AND aca.child = b.child", + " LEFT JOIN anon_parent_agg apa ON apa.parent = b.parent", + " CROSS JOIN anon_global_agg aga", + ] ) query_parts.append(")") @@ -400,8 +384,28 @@ async def _build_single_action_sql( restriction_intersect = "\nINTERSECT\n".join( f"SELECT * FROM ({sql})" for sql in restriction_sqls ) + # Decompose by NULL-pattern so the final filter can use pure-equality + # EXISTS lookups (satisfiable via automatic indexes) instead of a + # correlated OR-scan over the whole list. query_parts.extend( - [",", "restriction_list AS (", f" {restriction_intersect}", ")"] + [ + ",", + "restriction_list AS (", + f" {restriction_intersect}", + "),", + "restriction_exact AS (", + " SELECT parent, child FROM restriction_list WHERE parent IS NOT NULL AND child IS NOT NULL", + "),", + "restriction_parent_any AS (", + " SELECT DISTINCT parent FROM restriction_list WHERE parent IS NOT NULL AND child IS NULL", + "),", + "restriction_child_any AS (", + " SELECT DISTINCT child FROM restriction_list WHERE parent IS NULL AND child IS NOT NULL", + "),", + "restriction_all AS (", + " SELECT 1 AS matched FROM restriction_list WHERE parent IS NULL AND child IS NULL LIMIT 1", + ")", + ] ) # Final SELECT @@ -416,10 +420,11 @@ async def _build_single_action_sql( # Add restriction filter if there are restrictions if restriction_sqls: query_parts.append(""" - AND EXISTS ( - SELECT 1 FROM restriction_list r - WHERE (r.parent = decisions.parent OR r.parent IS NULL) - AND (r.child = decisions.child OR r.child IS NULL) + AND ( + EXISTS (SELECT 1 FROM restriction_all) + OR EXISTS (SELECT 1 FROM restriction_parent_any r WHERE r.parent = decisions.parent) + OR EXISTS (SELECT 1 FROM restriction_child_any r WHERE r.child = decisions.child) + OR EXISTS (SELECT 1 FROM restriction_exact r WHERE r.parent = decisions.parent AND r.child = decisions.child) )""") # Add parent filter if specified @@ -673,3 +678,239 @@ async def check_permission_for_resource( child=child, ) return results[action] + + +async def explain_permission_for_resource( + *, + datasette: "Datasette", + actor: dict | None, + action: str, + parent: str | None, + child: str | None, +) -> dict: + """Explain a permission decision for one action and resource. + + This is intended for Datasette's permission debugging tools. It uses the + same ``permission_resources_sql`` hook results and the same resolution + rules as :func:`check_permissions_for_actions`, but also returns the + matching rules, actor restriction results and ``also_requires`` chain. + + The returned dictionary is part of Datasette's unstable debugging API. + """ + + action_obj = datasette.actions.get(action) + if action_obj is None: + raise ValueError(f"Unknown action: {action}") + + explanation = await _explain_single_action( + datasette=datasette, + actor=actor, + action=action, + parent=parent, + child=child, + ) + + required_actions = [] + if action_obj.also_requires: + required = await explain_permission_for_resource( + datasette=datasette, + actor=actor, + action=action_obj.also_requires, + parent=parent, + child=child, + ) + required_actions.append(required) + + explanation["required_actions"] = required_actions + explanation["allowed"] = bool( + explanation["rule_allowed"] + and explanation["restriction_allowed"] + and all(required["allowed"] for required in required_actions) + ) + explanation["summary"] = _permission_explanation_summary(explanation) + return explanation + + +async def _explain_single_action( + *, + datasette: "Datasette", + actor: dict | None, + action: str, + parent: str | None, + child: str | None, +) -> dict: + """Return matching rules and restrictions for a single action.""" + from datasette.utils.permissions import SKIP_PERMISSION_CHECKS + + permission_sqls = await gather_permission_sql_from_hooks( + datasette=datasette, + actor=actor, + action=action, + ) + + if permission_sqls is SKIP_PERMISSION_CHECKS: + return { + "action": action, + "rule_allowed": True, + "restriction_allowed": True, + "winning_scope": "global", + "matched_rules": [ + { + "scope": "global", + "effect": "allow", + "source": "skip_permission_checks", + "reason": "Permission checks were explicitly skipped", + "decisive": True, + "ignored_because": None, + } + ], + "restrictions": [], + } + + db = datasette.get_internal_database() + matched_rules = [] + restrictions = [] + + for permission_sql in permission_sqls: + params = dict(permission_sql.params or {}) + parent_param = _unused_parameter_name(params, "_explain_parent") + params[parent_param] = parent + child_param = _unused_parameter_name(params, "_explain_child") + params[child_param] = child + + if permission_sql.sql: + rows = await db.execute( + f""" + SELECT parent, child, allow, reason + FROM ({permission_sql.sql}) AS permission_rules + WHERE (parent IS NULL OR parent = :{parent_param}) + AND (child IS NULL OR child = :{child_param}) + """, + params, + ) + for row in rows: + specificity = ( + 2 + if row["child"] is not None + else 1 if row["parent"] is not None else 0 + ) + matched_rules.append( + { + "scope": ("resource", "parent", "global")[2 - specificity], + "effect": "allow" if row["allow"] else "deny", + "source": permission_sql.source, + "reason": row["reason"], + "_specificity": specificity, + } + ) + + if permission_sql.restriction_sql: + restriction_row = ( + await db.execute( + f""" + SELECT EXISTS( + SELECT 1 FROM ({permission_sql.restriction_sql}) AS restriction_rules + WHERE (parent IS NULL OR parent = :{parent_param}) + AND (child IS NULL OR child = :{child_param}) + ) AS resource_is_in_allowlist + """, + params, + ) + ).first() + restriction_allowed = bool(restriction_row[0]) + restrictions.append( + { + "source": permission_sql.source, + "allowed": restriction_allowed, + "reason": params.get("deny") + or ( + "Resource is included in this restriction allowlist" + if restriction_allowed + else "Resource is not included in this restriction allowlist" + ), + } + ) + + matched_rules.sort( + key=lambda rule: ( + -rule["_specificity"], + 0 if rule["effect"] == "deny" else 1, + rule["source"] or "", + rule["reason"] or "", + ) + ) + + if matched_rules: + winning_specificity = matched_rules[0]["_specificity"] + winning_rules = [ + rule + for rule in matched_rules + if rule["_specificity"] == winning_specificity + ] + rule_allowed = not any(rule["effect"] == "deny" for rule in winning_rules) + winning_scope = winning_rules[0]["scope"] + else: + winning_specificity = None + rule_allowed = False + winning_scope = None + + for rule in matched_rules: + specificity = rule.pop("_specificity") + if specificity != winning_specificity: + rule["decisive"] = False + rule["ignored_because"] = "A more specific rule matched" + elif not rule_allowed and rule["effect"] == "allow": + rule["decisive"] = False + rule["ignored_because"] = "A deny rule matched at the same scope" + else: + rule["decisive"] = True + rule["ignored_because"] = None + + return { + "action": action, + "rule_allowed": rule_allowed, + "restriction_allowed": all( + restriction["allowed"] for restriction in restrictions + ), + "winning_scope": winning_scope, + "matched_rules": matched_rules, + "restrictions": restrictions, + } + + +def _unused_parameter_name(params: dict, preferred: str) -> str: + """Return a SQL parameter name that is not already in ``params``.""" + candidate = preferred + suffix = 2 + while candidate in params: + candidate = f"{preferred}_{suffix}" + suffix += 1 + return candidate + + +def _permission_explanation_summary(explanation: dict) -> str: + denied_requirement = next( + ( + required + for required in explanation["required_actions"] + if not required["allowed"] + ), + None, + ) + if denied_requirement: + return ( + f"Denied because {explanation['action']} also requires " + f"{denied_requirement['action']}, which was denied." + ) + if not explanation["matched_rules"]: + return "Denied because no permission rule matched this actor and resource." + if not explanation["rule_allowed"]: + return ( + f"Denied by a {explanation['winning_scope']}-level rule. " + "Deny rules take precedence over allow rules at the same scope." + ) + if not explanation["restriction_allowed"]: + return ( + "Denied because the resource is not included in the actor's restrictions." + ) + return f"Allowed by the matching {explanation['winning_scope']}-level rule." diff --git a/datasette/utils/asgi.py b/datasette/utils/asgi.py index 610b86f2..812194fd 100644 --- a/datasette/utils/asgi.py +++ b/datasette/utils/asgi.py @@ -1,28 +1,29 @@ import json -from typing import Optional +import re +from http.cookies import Morsel, SimpleCookie +from mimetypes import guess_type +from pathlib import Path +from urllib.parse import parse_qs, parse_qsl, urlunparse + +import aiofiles +import aiofiles.os + from datasette.utils import MultiParams, calculate_etag, error_body, sha256_file from datasette.utils.multipart import ( - parse_form_data, - MultipartParseError, - FormData, - DEFAULT_MAX_FILE_SIZE, - DEFAULT_MAX_REQUEST_SIZE, - DEFAULT_MAX_FIELDS, - DEFAULT_MAX_FILES, - DEFAULT_MAX_PARTS, DEFAULT_MAX_FIELD_SIZE, + DEFAULT_MAX_FIELDS, + DEFAULT_MAX_FILE_SIZE, + DEFAULT_MAX_FILES, DEFAULT_MAX_MEMORY_FILE_SIZE, DEFAULT_MAX_PART_HEADER_BYTES, DEFAULT_MAX_PART_HEADER_LINES, + DEFAULT_MAX_PARTS, + DEFAULT_MAX_REQUEST_SIZE, DEFAULT_MIN_FREE_DISK_BYTES, + FormData, + MultipartParseError, + parse_form_data, ) -from mimetypes import guess_type -from urllib.parse import parse_qs, urlunparse, parse_qsl -from pathlib import Path -from http.cookies import SimpleCookie, Morsel -import aiofiles -import aiofiles.os -import re # Workaround for adding samesite support to pre 3.8 python Morsel._reserved["samesite"] = "SameSite" @@ -88,7 +89,7 @@ class Request: self.max_post_body_bytes = max_post_body_bytes def __repr__(self): - return ''.format(self.method, self.url) + return f'' @property def method(self): @@ -167,7 +168,7 @@ class Request: if max_bytes is None: max_bytes = self.max_post_body_bytes too_large = PayloadTooLarge( - "Request body exceeded maximum size of {} bytes".format(max_bytes) + f"Request body exceeded maximum size of {max_bytes} bytes" ) if max_bytes: # Reject early if the client declares an oversized body @@ -206,7 +207,7 @@ class Request: max_request_size: int = DEFAULT_MAX_REQUEST_SIZE, max_fields: int = DEFAULT_MAX_FIELDS, max_files: int = DEFAULT_MAX_FILES, - max_parts: Optional[int] = DEFAULT_MAX_PARTS, + max_parts: int | None = DEFAULT_MAX_PARTS, max_field_size: int = DEFAULT_MAX_FIELD_SIZE, max_memory_file_size: int = DEFAULT_MAX_MEMORY_FILE_SIZE, max_part_header_bytes: int = DEFAULT_MAX_PART_HEADER_BYTES, @@ -529,9 +530,9 @@ class Response: httponly=False, samesite="lax", ): - assert samesite in SAMESITE_VALUES, "samesite should be one of {}".format( - SAMESITE_VALUES - ) + assert ( + samesite in SAMESITE_VALUES + ), f"samesite should be one of {SAMESITE_VALUES}" cookie = SimpleCookie() cookie[key] = value for prop_name, prop_value in ( diff --git a/datasette/utils/baseconv.py b/datasette/utils/baseconv.py index c4b64908..0469d7a8 100644 --- a/datasette/utils/baseconv.py +++ b/datasette/utils/baseconv.py @@ -13,7 +13,7 @@ Originally shared here: https://www.djangosnippets.org/snippets/1431/ """ -class BaseConverter(object): +class BaseConverter: decimal_digits = "0123456789" def __init__(self, digits): diff --git a/datasette/utils/check_callable.py b/datasette/utils/check_callable.py index a0997d20..e21a769b 100644 --- a/datasette/utils/check_callable.py +++ b/datasette/utils/check_callable.py @@ -1,6 +1,6 @@ import inspect import types -from typing import NamedTuple, Any +from typing import Any, NamedTuple class CallableStatus(NamedTuple): @@ -19,7 +19,7 @@ def check_callable(obj: Any) -> CallableStatus: if isinstance(obj, types.FunctionType): return CallableStatus(True, inspect.iscoroutinefunction(obj)) - if hasattr(obj, "__call__"): + if callable(obj): return CallableStatus(True, inspect.iscoroutinefunction(obj.__call__)) - assert False, "obj {} is somehow callable with no __call__ method".format(repr(obj)) + assert False, f"obj {obj!r} is somehow callable with no __call__ method" diff --git a/datasette/utils/internal_db.py b/datasette/utils/internal_db.py index 10ca32a5..0ddeb847 100644 --- a/datasette/utils/internal_db.py +++ b/datasette/utils/internal_db.py @@ -3,7 +3,7 @@ import textwrap from sqlite_utils import Database as SQLiteUtilsDatabase from sqlite_utils import Migrations -from datasette.utils import table_column_details +from datasette.utils import escape_sqlite, table_column_details INTERNAL_DB_SCHEMA_TABLES = { "catalog_databases", @@ -180,29 +180,9 @@ async def init_internal_db(db): await db.execute_write_fn(apply_migrations, transaction=False) -async def populate_schema_tables(internal_db, db): +async def populate_schema_tables(internal_db, db, schema_version): database_name = db.name - def delete_everything(conn): - conn.execute( - "DELETE FROM catalog_tables WHERE database_name = ?", [database_name] - ) - conn.execute( - "DELETE FROM catalog_views WHERE database_name = ?", [database_name] - ) - conn.execute( - "DELETE FROM catalog_columns WHERE database_name = ?", [database_name] - ) - conn.execute( - "DELETE FROM catalog_foreign_keys WHERE database_name = ?", - [database_name], - ) - conn.execute( - "DELETE FROM catalog_indexes WHERE database_name = ?", [database_name] - ) - - await internal_db.execute_write_fn(delete_everything) - tables = (await db.execute("select * from sqlite_master WHERE type = 'table'")).rows views = (await db.execute("select * from sqlite_master WHERE type = 'view'")).rows @@ -227,25 +207,30 @@ async def populate_schema_tables(internal_db, db): columns = table_column_details(conn, table_name) columns_to_insert.extend( { - **{"database_name": database_name, "table_name": table_name}, + "database_name": database_name, + "table_name": table_name, **column._asdict(), } for column in columns ) foreign_keys = conn.execute( - f"PRAGMA foreign_key_list([{table_name}])" + f"PRAGMA foreign_key_list({escape_sqlite(table_name)})" ).fetchall() foreign_keys_to_insert.extend( { - **{"database_name": database_name, "table_name": table_name}, + "database_name": database_name, + "table_name": table_name, **dict(foreign_key), } for foreign_key in foreign_keys ) - indexes = conn.execute(f"PRAGMA index_list([{table_name}])").fetchall() + indexes = conn.execute( + f"PRAGMA index_list({escape_sqlite(table_name)})" + ).fetchall() indexes_to_insert.extend( { - **{"database_name": database_name, "table_name": table_name}, + "database_name": database_name, + "table_name": table_name, **dict(index), } for index in indexes @@ -266,47 +251,76 @@ async def populate_schema_tables(internal_db, db): indexes_to_insert, ) = await db.execute_fn(collect_info) - await internal_db.execute_write_many( - """ - INSERT INTO catalog_tables (database_name, table_name, rootpage, sql) - values (?, ?, ?, ?) - """, - tables_to_insert, - ) - await internal_db.execute_write_many( - """ - INSERT INTO catalog_views (database_name, view_name, rootpage, sql) - values (?, ?, ?, ?) - """, - views_to_insert, - ) - await internal_db.execute_write_many( - """ - INSERT INTO catalog_columns ( - database_name, table_name, cid, name, type, "notnull", default_value, is_pk, hidden - ) VALUES ( - :database_name, :table_name, :cid, :name, :type, :notnull, :default_value, :is_pk, :hidden + def replace_catalog(conn): + # Delete child rows before their catalog_tables parents so this also + # works if a prepare_connection plugin enables foreign key enforcement. + for table in ( + "catalog_columns", + "catalog_foreign_keys", + "catalog_indexes", + "catalog_views", + "catalog_tables", + ): + conn.execute( + f"DELETE FROM {table} WHERE database_name = ?", + [database_name], + ) + conn.execute( + """ + INSERT OR REPLACE INTO catalog_databases ( + database_name, path, is_memory, schema_version + ) VALUES (?, ?, ?, ?) + """, + [ + database_name, + str(db.path) if db.path is not None else None, + db.is_memory, + schema_version, + ], ) - """, - columns_to_insert, - ) - await internal_db.execute_write_many( - """ - INSERT INTO catalog_foreign_keys ( - database_name, table_name, "id", seq, "table", "from", "to", on_update, on_delete, match - ) VALUES ( - :database_name, :table_name, :id, :seq, :table, :from, :to, :on_update, :on_delete, :match + conn.executemany( + """ + INSERT INTO catalog_tables (database_name, table_name, rootpage, sql) + values (?, ?, ?, ?) + """, + tables_to_insert, ) - """, - foreign_keys_to_insert, - ) - await internal_db.execute_write_many( - """ - INSERT INTO catalog_indexes ( - database_name, table_name, seq, name, "unique", origin, partial - ) VALUES ( - :database_name, :table_name, :seq, :name, :unique, :origin, :partial + conn.executemany( + """ + INSERT INTO catalog_views (database_name, view_name, rootpage, sql) + values (?, ?, ?, ?) + """, + views_to_insert, ) - """, - indexes_to_insert, - ) + conn.executemany( + """ + INSERT INTO catalog_columns ( + database_name, table_name, cid, name, type, "notnull", default_value, is_pk, hidden + ) VALUES ( + :database_name, :table_name, :cid, :name, :type, :notnull, :default_value, :is_pk, :hidden + ) + """, + columns_to_insert, + ) + conn.executemany( + """ + INSERT INTO catalog_foreign_keys ( + database_name, table_name, "id", seq, "table", "from", "to", on_update, on_delete, match + ) VALUES ( + :database_name, :table_name, :id, :seq, :table, :from, :to, :on_update, :on_delete, :match + ) + """, + foreign_keys_to_insert, + ) + conn.executemany( + """ + INSERT INTO catalog_indexes ( + database_name, table_name, seq, name, "unique", origin, partial + ) VALUES ( + :database_name, :table_name, :seq, :name, :unique, :origin, :partial + ) + """, + indexes_to_insert, + ) + + await internal_db.execute_write_fn(replace_catalog) diff --git a/datasette/utils/multipart.py b/datasette/utils/multipart.py index cfa77486..13289b1c 100644 --- a/datasette/utils/multipart.py +++ b/datasette/utils/multipart.py @@ -11,15 +11,10 @@ Supports: import asyncio import shutil import tempfile +from collections.abc import Callable from dataclasses import dataclass, field from typing import ( Any, - Callable, - Dict, - List, - Optional, - Tuple, - Union, ) from urllib.parse import parse_qsl @@ -29,7 +24,7 @@ DEFAULT_MAX_REQUEST_SIZE = 100 * 1024 * 1024 # 100MB DEFAULT_MAX_FIELDS = 1000 DEFAULT_MAX_FILES = 100 # If max_parts is not specified, it defaults to max_fields + max_files -DEFAULT_MAX_PARTS: Optional[int] = None +DEFAULT_MAX_PARTS: int | None = None DEFAULT_MAX_FIELD_SIZE = 100 * 1024 # 100KB DEFAULT_MAX_MEMORY_FILE_SIZE = 1024 * 1024 # 1MB DEFAULT_MAX_PART_HEADER_BYTES = 16 * 1024 # 16KB @@ -40,8 +35,6 @@ DEFAULT_MIN_FREE_DISK_BYTES = 50 * 1024 * 1024 # 50MB class MultipartParseError(Exception): """Raised when multipart parsing fails.""" - pass - @dataclass class UploadedFile: @@ -57,7 +50,7 @@ class UploadedFile: name: str filename: str - content_type: Optional[str] + content_type: str | None size: int _file: tempfile.SpooledTemporaryFile = field(repr=False) @@ -86,7 +79,8 @@ class UploadedFile: def __del__(self): try: self._file.close() - except Exception: + except Exception: # noqa: BLE001, S110 + # __del__ must never raise pass @@ -98,27 +92,27 @@ class FormData: """ def __init__(self): - self._data: List[Tuple[str, Union[str, UploadedFile]]] = [] + self._data: list[tuple[str, str | UploadedFile]] = [] - def append(self, key: str, value: Union[str, UploadedFile]) -> None: + def append(self, key: str, value: str | UploadedFile) -> None: """Add a key-value pair.""" self._data.append((key, value)) - def __getitem__(self, key: str) -> Union[str, UploadedFile]: + def __getitem__(self, key: str) -> str | UploadedFile: """Get the first value for a key.""" for k, v in self._data: if k == key: return v raise KeyError(key) - def get(self, key: str, default: Any = None) -> Optional[Union[str, UploadedFile]]: + def get(self, key: str, default: Any = None) -> str | UploadedFile | None: """Get the first value for a key, or default if not found.""" try: return self[key] except KeyError: return default - def getlist(self, key: str) -> List[Union[str, UploadedFile]]: + def getlist(self, key: str) -> list[str | UploadedFile]: """Get all values for a key.""" return [v for k, v in self._data if k == key] @@ -142,15 +136,15 @@ class FormData: """Return unique keys.""" return list(self) - def items(self) -> List[Tuple[str, Union[str, UploadedFile]]]: + def items(self) -> list[tuple[str, str | UploadedFile]]: """Return all key-value pairs.""" return list(self._data) - def values(self) -> List[Union[str, UploadedFile]]: + def values(self) -> list[str | UploadedFile]: """Return all values.""" return [v for _, v in self._data] - def _uploaded_files(self) -> List[UploadedFile]: + def _uploaded_files(self) -> list[UploadedFile]: """Return UploadedFile instances contained in this form.""" return [v for _, v in self._data if isinstance(v, UploadedFile)] @@ -163,7 +157,7 @@ class FormData: for uploaded in self._uploaded_files(): try: uploaded.close_sync() - except Exception: + except Exception: # noqa: BLE001, S110 # Best-effort cleanup; ignore close errors pass @@ -172,7 +166,7 @@ class FormData: for uploaded in self._uploaded_files(): try: await uploaded.close() - except Exception: + except Exception: # noqa: BLE001, S110 # Best-effort cleanup; ignore close errors pass @@ -189,13 +183,13 @@ class FormData: await self.aclose() -def parse_content_disposition(header: str) -> Dict[str, Optional[str]]: +def parse_content_disposition(header: str) -> dict[str, str | None]: """ Parse Content-Disposition header value. Returns dict with 'name', 'filename' keys (filename may be None). """ - result: Dict[str, Optional[str]] = {"name": None, "filename": None} + result: dict[str, str | None] = {"name": None, "filename": None} # Split on semicolons, handling quoted strings parts = [] @@ -238,7 +232,8 @@ def parse_content_disposition(header: str) -> Dict[str, Optional[str]]: from urllib.parse import unquote result["filename"] = unquote(encoded, encoding="utf-8") - except Exception: + except Exception: # noqa: BLE001, S110 + # Malformed RFC 5987 filename* - fall back to the plain filename pass continue @@ -250,20 +245,19 @@ def parse_content_disposition(header: str) -> Dict[str, Optional[str]]: if key == "name": result["name"] = value - elif key == "filename": - # Only set if filename* hasn't already set it - if result["filename"] is None: - # Strip path components (security) - # Handle both Unix and Windows paths - value = value.replace("\\", "/") - if "/" in value: - value = value.rsplit("/", 1)[-1] - result["filename"] = value + # Only set filename if filename* hasn't already set it + elif key == "filename" and result["filename"] is None: + # Strip path components (security) + # Handle both Unix and Windows paths + value = value.replace("\\", "/") + if "/" in value: + value = value.rsplit("/", 1)[-1] + result["filename"] = value return result -def parse_content_type(header: str) -> Tuple[str, Dict[str, str]]: +def parse_content_type(header: str) -> tuple[str, dict[str, str]]: """ Parse Content-Type header value. @@ -307,7 +301,7 @@ class MultipartParser: max_request_size: int = DEFAULT_MAX_REQUEST_SIZE, max_fields: int = DEFAULT_MAX_FIELDS, max_files: int = DEFAULT_MAX_FILES, - max_parts: Optional[int] = DEFAULT_MAX_PARTS, + max_parts: int | None = DEFAULT_MAX_PARTS, max_field_size: int = DEFAULT_MAX_FIELD_SIZE, max_memory_file_size: int = DEFAULT_MAX_MEMORY_FILE_SIZE, max_part_header_bytes: int = DEFAULT_MAX_PART_HEADER_BYTES, @@ -348,12 +342,12 @@ class MultipartParser: self._tempdir = tempfile.gettempdir() # Current part state - self.current_headers: Dict[str, str] = {} - self.current_file: Optional[tempfile.SpooledTemporaryFile] = None + self.current_headers: dict[str, str] = {} + self.current_file: tempfile.SpooledTemporaryFile | None = None self.current_body = bytearray() - self.current_name: Optional[str] = None - self.current_filename: Optional[str] = None - self.current_content_type: Optional[str] = None + self.current_name: str | None = None + self.current_filename: str | None = None + self.current_content_type: str | None = None def feed(self, chunk: bytes) -> None: """Feed a chunk of data to the parser.""" @@ -454,7 +448,7 @@ class MultipartParser: # Parse header try: line_str = line.decode("utf-8", errors="replace") - except Exception: + except UnicodeDecodeError: line_str = line.decode("latin-1") if ":" in line_str: @@ -481,7 +475,9 @@ class MultipartParser: if self.file_count > self.max_files: raise MultipartParseError("Too many files") if self.handle_files: - self.current_file = tempfile.SpooledTemporaryFile( + # Outlives this method - it is filled in across parser callbacks + # and then handed to the UploadedFile the caller consumes + self.current_file = tempfile.SpooledTemporaryFile( # noqa: SIM115 max_size=self.max_memory_file_size ) else: @@ -644,7 +640,7 @@ async def parse_form_data( max_request_size: int = DEFAULT_MAX_REQUEST_SIZE, max_fields: int = DEFAULT_MAX_FIELDS, max_files: int = DEFAULT_MAX_FILES, - max_parts: Optional[int] = DEFAULT_MAX_PARTS, + max_parts: int | None = DEFAULT_MAX_PARTS, max_field_size: int = DEFAULT_MAX_FIELD_SIZE, max_memory_file_size: int = DEFAULT_MAX_MEMORY_FILE_SIZE, max_part_header_bytes: int = DEFAULT_MAX_PART_HEADER_BYTES, diff --git a/datasette/utils/permissions.py b/datasette/utils/permissions.py index fd1e41a1..5a8ee8e2 100644 --- a/datasette/utils/permissions.py +++ b/datasette/utils/permissions.py @@ -2,8 +2,9 @@ from __future__ import annotations import json -from typing import Any, Dict, Iterable, List, Sequence, Tuple import sqlite3 +from collections.abc import Iterable, Sequence +from typing import Any from datasette.permissions import PermissionSQL from datasette.plugins import pm @@ -15,7 +16,7 @@ SKIP_PERMISSION_CHECKS = object() async def gather_permission_sql_from_hooks( *, datasette, actor: dict | None, action: str -) -> List[PermissionSQL] | object: +) -> list[PermissionSQL] | object: """Collect PermissionSQL objects from the permission_resources_sql hook. Ensures that each returned PermissionSQL has a populated ``source``. @@ -34,7 +35,7 @@ async def gather_permission_sql_from_hooks( hookimpls = hook_caller.get_hookimpls() hook_results = list(hook_caller(datasette=datasette, actor=actor, action=action)) - collected: List[PermissionSQL] = [] + collected: list[PermissionSQL] = [] actor_json = json.dumps(actor) if actor is not None else None actor_id = actor.get("id") if isinstance(actor, dict) else None @@ -71,7 +72,7 @@ def _iter_permission_sql_from_result( if isinstance(result, PermissionSQL): return [result] if isinstance(result, (list, tuple)): - collected: List[PermissionSQL] = [] + collected: list[PermissionSQL] = [] for item in result: collected.extend(_iter_permission_sql_from_result(item, action=action)) return collected @@ -90,7 +91,7 @@ def _iter_permission_sql_from_result( def build_rules_union( actor: dict | None, plugins: Sequence[PermissionSQL] -) -> Tuple[str, Dict[str, Any]]: +) -> tuple[str, dict[str, Any]]: """ Compose plugin SQL into a UNION ALL. @@ -102,10 +103,10 @@ def build_rules_union( The system reserves these parameter names: :actor, :actor_id, :action, :filter_parent Plugin parameters should be prefixed with a unique identifier (e.g., source name). """ - parts: List[str] = [] + parts: list[str] = [] actor_json = json.dumps(actor) if actor else None actor_id = actor.get("id") if actor else None - params: Dict[str, Any] = {"actor": actor_json, "actor_id": actor_id} + params: dict[str, Any] = {"actor": actor_json, "actor_id": actor_id} for p in plugins: # No namespacing - just use plugin params as-is @@ -141,10 +142,10 @@ async def resolve_permissions_from_catalog( plugins: Sequence[Any], action: str, candidate_sql: str, - candidate_params: Dict[str, Any] | None = None, + candidate_params: dict[str, Any] | None = None, *, implicit_deny: bool = True, -) -> List[Dict[str, Any]]: +) -> list[dict[str, Any]]: """ Resolve permissions by embedding the provided *candidate_sql* in a CTE. @@ -168,8 +169,8 @@ async def resolve_permissions_from_catalog( - parent, child, allow, reason, source_plugin, depth - resource (rendered "/parent/child" or "/parent" or "/") """ - resolved_plugins: List[PermissionSQL] = [] - restriction_sqls: List[str] = [] + resolved_plugins: list[PermissionSQL] = [] + restriction_sqls: list[str] = [] for plugin in plugins: if callable(plugin) and not isinstance(plugin, PermissionSQL): @@ -398,11 +399,11 @@ async def resolve_permissions_with_candidates( db, actor: dict | None, plugins: Sequence[Any], - candidates: List[Tuple[str, str | None]], + candidates: list[tuple[str, str | None]], action: str, *, implicit_deny: bool = True, -) -> List[Dict[str, Any]]: +) -> list[dict[str, Any]]: """ Resolve permissions without any external candidate table by embedding the candidates as a UNION of parameterized SELECTs in a CTE. @@ -411,8 +412,8 @@ async def resolve_permissions_with_candidates( actor: actor dict (or None), made available as :actor (JSON), :actor_id, and :action """ # Build a small CTE for candidates. - cand_rows_sql: List[str] = [] - cand_params: Dict[str, Any] = {} + cand_rows_sql: list[str] = [] + cand_params: dict[str, Any] = {} for i, (parent, child) in enumerate(candidates): pkey = f"cand_p_{i}" ckey = f"cand_c_{i}" diff --git a/datasette/utils/shutil_backport.py b/datasette/utils/shutil_backport.py index d1fd1bd7..d323f5d6 100644 --- a/datasette/utils/shutil_backport.py +++ b/datasette/utils/shutil_backport.py @@ -6,7 +6,7 @@ https://github.com/python/cpython/blob/v3.8.3/LICENSE """ import os -from shutil import copy, copy2, copystat, Error +from shutil import Error, copy, copy2, copystat def _copytree( diff --git a/datasette/utils/sql_analysis.py b/datasette/utils/sql_analysis.py index 1be28982..334545bd 100644 --- a/datasette/utils/sql_analysis.py +++ b/datasette/utils/sql_analysis.py @@ -413,12 +413,12 @@ def analyze_sql_tables( database=None, table=None, sqlite_schema=sqlite_schema, - target="{} {}".format(arg1, arg2) if arg2 is not None else arg1, + target=f"{arg1} {arg2}" if arg2 is not None else arg1, source=source, ) return sqlite3.SQLITE_OK - action_name = _AUTHORIZER_ACTION_NAMES.get(action, "SQLITE_{}".format(action)) + action_name = _AUTHORIZER_ACTION_NAMES.get(action, f"SQLITE_{action}") record( "unknown", "unknown", @@ -521,9 +521,7 @@ def analyze_sql_tables( and key.target in _SQLITE_INTERNAL_SCHEMA_FUNCTIONS ): return True - if key_is_drop_table_delete(key): - return True - return False + return bool(key_is_drop_table_delete(key)) def table_kind_for(key: OperationKey) -> SQLiteTableType | None: if ( diff --git a/datasette/utils/sqlite.py b/datasette/utils/sqlite.py index 4743ae4c..d3926f6f 100644 --- a/datasette/utils/sqlite.py +++ b/datasette/utils/sqlite.py @@ -100,7 +100,7 @@ def sqlite_hidden_table_names(conn, *, schema: str | None = "main") -> list[str] schema_table = _sqlite_schema_table(schema) try: rows = conn.execute( - "select name, sql from {} where type = 'table'".format(schema_table) + f"select name, sql from {schema_table} where type = 'table'" ).fetchall() except sqlite3.DatabaseError: return [] @@ -127,7 +127,7 @@ def _sqlite_table_type_from_schema( schema_table = _sqlite_schema_table(schema) try: row = conn.execute( - "select type, sql from {} where name = ?".format(schema_table), + f"select type, sql from {schema_table} where name = ?", (table,), ).fetchone() except sqlite3.DatabaseError: @@ -155,7 +155,7 @@ def _is_known_shadow_table( schema_table = _sqlite_schema_table(schema) try: rows = conn.execute( - "select name, sql from {} where type = 'table'".format(schema_table) + f"select name, sql from {schema_table} where type = 'table'" ).fetchall() except sqlite3.DatabaseError: return False @@ -174,7 +174,7 @@ def _sqlite_schema_table(schema: str | None) -> str: return "sqlite_master" if schema == "temp": return "sqlite_temp_master" - return "{}.sqlite_master".format(_quote_identifier(schema)) + return f"{_quote_identifier(schema)}.sqlite_master" def _quote_identifier(value: str) -> str: diff --git a/datasette/utils/testing.py b/datasette/utils/testing.py index de7e94af..a8be47bf 100644 --- a/datasette/utils/testing.py +++ b/datasette/utils/testing.py @@ -1,6 +1,7 @@ -from asgiref.sync import async_to_sync -from urllib.parse import urlencode import json +from urllib.parse import urlencode + +from asgiref.sync import async_to_sync # These wrapper classes pre-date the introduction of # datasette.client and httpx to Datasette. They could diff --git a/datasette/version.py b/datasette/version.py index 387144e9..8e238ab5 100644 --- a/datasette/version.py +++ b/datasette/version.py @@ -1,2 +1,2 @@ -__version__ = "1.0a36" +__version__ = "1.0a37" __version_info__ = tuple(__version__.split(".")) diff --git a/datasette/views/__init__.py b/datasette/views/__init__.py index ed7e175f..bac3b39e 100644 --- a/datasette/views/__init__.py +++ b/datasette/views/__init__.py @@ -1,7 +1,7 @@ -from dataclasses import dataclass import dataclasses import types import typing +from dataclasses import dataclass @dataclass(frozen=True) @@ -74,16 +74,14 @@ class Context: extra_class = table_extra_registry.classes_by_name[name] except KeyError: raise KeyError( - "{}.{} is declared with from_extra() but there is no " - "registered extra of that name".format(cls.__name__, name) + f"{cls.__name__}.{name} is declared with from_extra() but there is no " + "registered extra of that name" ) if cls.extras_scope is not None and not extra_class.available_for( cls.extras_scope ): raise ValueError( - "{}.{} is declared with from_extra() but the {} extra is " - "not available for scope {}".format( - cls.__name__, name, name, cls.extras_scope - ) + f"{cls.__name__}.{name} is declared with from_extra() but the {name} extra is " + f"not available for scope {cls.extras_scope}" ) return extra_class.description or "" diff --git a/datasette/views/base.py b/datasette/views/base.py index 66e14a6d..48108cec 100644 --- a/datasette/views/base.py +++ b/datasette/views/base.py @@ -2,20 +2,20 @@ import csv import hashlib import sys -from datasette.utils.asgi import Request from datasette.utils import ( - add_cors_headers, EscapeHtmlWriter, InvalidSql, LimitedWriter, + add_cors_headers, path_from_row_pks, path_with_format, sqlite3, ) from datasette.utils.asgi import ( AsgiStream, - Response, BadRequest, + Request, + Response, ) @@ -129,12 +129,10 @@ class BaseView: template = environment.select_template(templates) template_context = { **context, - **{ - "select_templates": [ - f"{'*' if template_name == template.name else ''}{template_name}" - for template_name in templates - ], - }, + "select_templates": [ + f"{'*' if template_name == template.name else ''}{template_name}" + for template_name in templates + ], } headers = {} if self.has_json_alternate: @@ -151,9 +149,7 @@ class BaseView: template_context["alternate_url_json"] = alternate_url_json headers.update( { - "Link": '<{}>; rel="alternate"; type="application/json+datasette"'.format( - alternate_url_json - ) + "Link": f'<{alternate_url_json}>; rel="alternate"; type="application/json+datasette"' } ) return Response.html( @@ -184,9 +180,7 @@ async def stream_csv(datasette, fetch_data, request, database): stream = request.args.get("_stream") # Do not calculate facets or counts: extra_parameters = [ - "{}=1".format(key) - for key in ("_nofacet", "_nocount") - if not request.args.get(key) + f"{key}=1" for key in ("_nofacet", "_nocount") if not request.args.get(key) ] if extra_parameters: # Replace request object with a new one with modified scope @@ -216,9 +210,6 @@ async def stream_csv(datasette, fetch_data, request, database): except (sqlite3.OperationalError, InvalidSql) as e: raise DatasetteError(str(e), title="Invalid SQL", status=400) - except sqlite3.OperationalError as e: - raise DatasetteError(str(e)) - except DatasetteError: raise @@ -325,8 +316,9 @@ async def stream_csv(datasette, fetch_data, request, database): else: new_row.append(cell) await writer.writerow(new_row) - except Exception as ex: - sys.stderr.write("Caught this error: {}\n".format(ex)) + except Exception as ex: # noqa: BLE001 + # Streaming CSV: report the error into the response body and stop + sys.stderr.write(f"Caught this error: {ex}\n") sys.stderr.flush() await r.write(str(ex)) return diff --git a/datasette/views/database.py b/datasette/views/database.py index 10dc66ae..f54ffd38 100644 --- a/datasette/views/database.py +++ b/datasette/views/database.py @@ -1,49 +1,52 @@ -from dataclasses import asdict, dataclass, field -from urllib.parse import parse_qsl, urlencode import asyncio import hashlib import itertools import json -import markupsafe import os import textwrap +from dataclasses import asdict, dataclass, field +from urllib.parse import parse_qsl, urlencode + +import markupsafe -from datasette.extras import extra_names_from_request, ExtraScope from datasette.database import QueryInterrupted +from datasette.extras import ExtraScope, extra_names_from_request +from datasette.plugins import pm from datasette.resources import DatabaseResource, QueryResource from datasette.stored_queries import StoredQuery, stored_query_to_dict -from datasette.write_sql import QueryWriteRejected from datasette.utils import ( + InvalidSql, add_cors_headers, await_me_maybe, - error_body, call_with_supported_arguments, - named_parameters as derive_named_parameters, + error_body, format_bytes, - make_slot_function, - tilde_decode, - to_css_class, - validate_sql_select, is_url, + make_slot_function, path_with_added_args, path_with_format, path_with_removed_args, sqlite3, + tilde_decode, + to_css_class, truncate_url, - InvalidSql, + validate_sql_select, ) -from datasette.utils.asgi import AsgiFileDownload, NotFound, Response, Forbidden -from datasette.plugins import pm +from datasette.utils import ( + named_parameters as derive_named_parameters, +) +from datasette.utils.asgi import AsgiFileDownload, Forbidden, NotFound, Response +from datasette.write_sql import QueryWriteRejected +from . import Context from .base import DatasetteError, View, stream_csv from .query_helpers import _ensure_stored_query_execution_permissions, _table_columns +from .table_create_alter import _create_table_ui_context from .table_extras import ( QueryExtraContext, resolve_query_extras, table_extra_registry, ) -from .table_create_alter import _create_table_ui_context -from . import Context @dataclass @@ -100,7 +103,7 @@ class DatabaseView(View): return response if format_ not in ("html", "json"): - raise NotFound("Invalid format: {}".format(format_)) + raise NotFound(f"Invalid format: {format_}") metadata = await datasette.get_database_metadata(database) @@ -164,7 +167,7 @@ class DatabaseView(View): "label": "Create table", "description": "Create a new table in this database.", "attrs": { - "aria-label": "Create table in {}".format(database), + "aria-label": f"Create table in {database}", "data-database-action": "create-table", }, } @@ -271,9 +274,7 @@ class DatabaseView(View): view_name="database", ), headers={ - "Link": '<{}>; rel="alternate"; type="application/json+datasette"'.format( - alternate_url_json - ) + "Link": f'<{alternate_url_json}>; rel="alternate"; type="application/json+datasette"' }, ) @@ -556,7 +557,7 @@ async def database_download(request, datasette): if datasette.cors: add_cors_headers(headers) if db.hash: - etag = '"{}"'.format(db.hash) + etag = f'"{db.hash}"' headers["Etag"] = etag # Has user seen this already? if_none_match = request.headers.get("if-none-match") @@ -643,8 +644,15 @@ class QueryView(View): ok = None redirect_url = None try: + execute_write_kwargs = {"request": request} + if stored_query.is_trusted: + analysis = await db.analyze_sql(stored_query.sql, params_for_query) + if any( + operation.operation == "vacuum" for operation in analysis.operations + ): + execute_write_kwargs["transaction"] = False cursor = await db.execute_write( - stored_query.sql, params_for_query, request=request + stored_query.sql, params_for_query, **execute_write_kwargs ) # success message can come from on_success_message or on_success_message_sql message = None @@ -657,8 +665,9 @@ class QueryView(View): ).first() if message_result: message = message_result[0] - except Exception as ex: - message = "Error running on_success_message_sql: {}".format(ex) + except Exception as ex: # noqa: BLE001 + # Stored-query on_success_message_sql is user-authored + message = f"Error running on_success_message_sql: {ex}" message_type = datasette.ERROR if not message: if stored_query.on_success_message: @@ -672,7 +681,8 @@ class QueryView(View): redirect_url = stored_query.on_success_redirect ok = True - except Exception as ex: + except Exception as ex: # noqa: BLE001 + # Stored-query execution is user-authored SQL message = stored_query.on_error_message or str(ex) message_type = datasette.ERROR redirect_url = stored_query.on_error_redirect @@ -806,16 +816,16 @@ class QueryView(View): rows = results.rows except QueryInterrupted as ex: raise DatasetteError( - textwrap.dedent(""" + textwrap.dedent(f"""

SQL query took too long. The time limit is controlled by the sql_time_limit_ms configuration option.

- + - """.format(markupsafe.escape(ex.sql))).strip(), + """).strip(), title="SQL Interrupted", status=400, message_is_html=True, @@ -831,8 +841,6 @@ class QueryView(View): columns = [] except (sqlite3.OperationalError, InvalidSql) as ex: raise DatasetteError(str(ex), title="Invalid SQL", status=400) - except sqlite3.OperationalError as ex: - raise DatasetteError(str(ex)) except DatasetteError: raise @@ -854,7 +862,7 @@ class QueryView(View): return data, None, None return await stream_csv(datasette, fetch_data_for_csv, request, db.name) - elif format_ in datasette.renderers.keys(): + elif format_ in datasette.renderers: if not sql: raise DatasetteError("?sql= is required", status=400) data = {"ok": True, "rows": rows, "columns": columns} @@ -946,9 +954,7 @@ class QueryView(View): } headers.update( { - "Link": '<{}>; rel="alternate"; type="application/json+datasette"'.format( - alternate_url_json - ) + "Link": f'<{alternate_url_json}>; rel="alternate"; type="application/json+datasette"' } ) metadata = await query_metadata() @@ -1029,9 +1035,7 @@ class QueryView(View): + "?" + urlencode( { - **{ - "sql": sql, - }, + "sql": sql, **named_parameter_values, } ) @@ -1133,7 +1137,7 @@ class QueryView(View): headers=headers, ) else: - assert False, "Invalid format: {}".format(format_) + assert False, f"Invalid format: {format_}" if datasette.cors: add_cors_headers(r.headers) return r @@ -1234,7 +1238,7 @@ async def display_rows(datasette, database, request, rows, columns): '<Binary: {:,} byte{}>'.format( blob_url, ( - ' title="{}"'.format(formatted) + f' title="{formatted}"' if "bytes" not in formatted else "" ), diff --git a/datasette/views/execute_write.py b/datasette/views/execute_write.py index dd35b127..3dfb810c 100644 --- a/datasette/views/execute_write.py +++ b/datasette/views/execute_write.py @@ -8,8 +8,8 @@ from datasette.utils.asgi import Response from .base import BaseView from .database import display_rows as display_query_rows from .query_helpers import ( - QueryValidationError, SQL_PARAMETER_FORM_PREFIX, + QueryValidationError, _analysis_is_write, _analysis_rows, _analysis_rows_with_permissions, @@ -31,15 +31,7 @@ WRITE_TEMPLATE_LABELS = { "delete": "Delete rows", } WRITE_TEMPLATE_OPERATIONS = tuple(WRITE_TEMPLATE_LABELS) -CREATE_TABLE_TEMPLATE_SQL = "\n".join( - ( - "create table new_table (", - " id integer primary key,", - " name text", - " -- created text default (datetime('now'))", - ")", - ) -) +CREATE_TABLE_TEMPLATE_SQL = "create table new_table (\n id integer primary key,\n name text\n -- created text default (datetime('now'))\n)" def _parameter_names(columns): @@ -49,11 +41,11 @@ def _parameter_names(columns): base = re.sub(r"[^a-z0-9_]+", "_", column.lower()) base = base.strip("_") or "value" if base[0].isdigit(): - base = "p_{}".format(base) + base = f"p_{base}" name = base index = 2 while name in seen: - name = "{}_{}".format(base, index) + name = f"{base}_{index}" index += 1 seen.add(name) names[column] = name @@ -65,7 +57,7 @@ def _quote_identifier(identifier): def _preferred_where_column(table, columns): - lower_table_id = "{}_id".format(table.lower()) + lower_table_id = f"{table.lower()}_id" return ( next((column for column in columns if column.lower() == "id"), None) or next( @@ -90,17 +82,15 @@ def _insert_template_sql(table, columns): auto_pk = _auto_incrementing_primary_key(columns) insert_columns = [column for column in column_names if column != auto_pk] if not insert_columns: - return "insert into {}\ndefault values".format(_quote_identifier(table)) + return f"insert into {_quote_identifier(table)}\ndefault values" names = _parameter_names(insert_columns) return "\n".join( ( - "insert into {} (".format(_quote_identifier(table)), - ",\n".join( - " {}".format(_quote_identifier(column)) for column in insert_columns - ), + f"insert into {_quote_identifier(table)} (", + ",\n".join(f" {_quote_identifier(column)}" for column in insert_columns), ")", "values (", - ",\n".join(" :{}".format(names[column]) for column in insert_columns), + ",\n".join(f" :{names[column]}" for column in insert_columns), ")", ) ) @@ -114,18 +104,14 @@ def _update_template_sql(table, columns): if not set_columns: return "\n".join( ( - "update {}".format(_quote_identifier(table)), - "set {} = :new_{}".format( - _quote_identifier(where_column), names[where_column] - ), - "where {} = :{}".format( - _quote_identifier(where_column), names[where_column] - ), + f"update {_quote_identifier(table)}", + f"set {_quote_identifier(where_column)} = :new_{names[where_column]}", + f"where {_quote_identifier(where_column)} = :{names[where_column]}", ) ) return "\n".join( ( - "update {}".format(_quote_identifier(table)), + f"update {_quote_identifier(table)}", "set " + ",\n".join( "{}{} = :{}".format( @@ -135,9 +121,7 @@ def _update_template_sql(table, columns): ) for index, column in enumerate(set_columns) ), - "where {} = :{}".format( - _quote_identifier(where_column), names[where_column] - ), + f"where {_quote_identifier(where_column)} = :{names[where_column]}", ) ) @@ -148,10 +132,8 @@ def _delete_template_sql(table, columns): where_column = _preferred_where_column(table, column_names) return "\n".join( ( - "delete from {}".format(_quote_identifier(table)), - "where {} = :{}".format( - _quote_identifier(where_column), names[where_column] - ), + f"delete from {_quote_identifier(table)}", + f"where {_quote_identifier(where_column)} = :{names[where_column]}", ) ) diff --git a/datasette/views/index.py b/datasette/views/index.py index 67296cd1..f73ee38a 100644 --- a/datasette/views/index.py +++ b/datasette/views/index.py @@ -2,11 +2,11 @@ import json from datasette.plugins import pm from datasette.utils import ( + UNSTABLE_API_MESSAGE, + CustomJSONEncoder, add_cors_headers, await_me_maybe, make_slot_function, - CustomJSONEncoder, - UNSTABLE_API_MESSAGE, ) from datasette.utils.asgi import Response from datasette.version import __version__ @@ -46,15 +46,15 @@ class IndexView(BaseView): databases = [] # Iterate over allowed databases instead of all databases - for name in allowed_db_dict.keys(): + for name, allowed_db in allowed_db_dict.items(): db = self.ds.databases[name] - database_private = allowed_db_dict[name].private + database_private = allowed_db.private # Get allowed tables/views for this database allowed_for_db = tables_by_db.get(name, {}) # Get table names from allowed set instead of db.table_names() - table_names = [child_name for child_name in allowed_for_db.keys()] + table_names = [child_name for child_name in allowed_for_db] hidden_table_names = set(await db.hidden_table_names()) @@ -99,7 +99,7 @@ class IndexView(BaseView): # We will be sorting by number of relationships, so populate that field all_foreign_keys = await db.get_all_foreign_keys() for table, foreign_keys in all_foreign_keys.items(): - if table in tables.keys(): + if table in tables: count = len(foreign_keys["incoming"] + foreign_keys["outgoing"]) tables[table]["num_relationships_for_sorting"] = count @@ -121,8 +121,7 @@ class IndexView(BaseView): # Only add views if this is less than TRUNCATE_AT if len(tables_and_views_truncated) < TRUNCATE_AT: num_views_to_add = TRUNCATE_AT - len(tables_and_views_truncated) - for view in views[:num_views_to_add]: - tables_and_views_truncated.append(view) + tables_and_views_truncated.extend(views[:num_views_to_add]) databases.append( { diff --git a/datasette/views/query_helpers.py b/datasette/views/query_helpers.py index 588891d4..725d9cdb 100644 --- a/datasette/views/query_helpers.py +++ b/datasette/views/query_helpers.py @@ -5,6 +5,19 @@ from datasette.resources import DatabaseResource from datasette.stored_queries import ( StoredQuery, ) +from datasette.utils import ( + InvalidSql, + escape_sqlite, + parse_size_limit, + path_from_row_pks, + sqlite3, + validate_sql_select, +) +from datasette.utils import ( + named_parameters as derive_named_parameters, +) +from datasette.utils.asgi import Forbidden +from datasette.utils.sql_analysis import Operation, SQLAnalysis from datasette.write_sql import ( IgnoreWriteSqlOperation, QueryWriteRejected, @@ -12,17 +25,6 @@ from datasette.write_sql import ( decision_for_write_sql_operation, operation_is_write, ) -from datasette.utils import ( - parse_size_limit, - named_parameters as derive_named_parameters, - escape_sqlite, - path_from_row_pks, - sqlite3, - validate_sql_select, - InvalidSql, -) -from datasette.utils.asgi import Forbidden -from datasette.utils.sql_analysis import Operation, SQLAnalysis _query_name_re = re.compile(r"^[^/\.\n]+$") @@ -91,7 +93,7 @@ def _as_optional_bool(value, name): return True if lowered in {"0", "false", "f", "no", "off"}: return False - raise QueryValidationError("{} must be 0 or 1".format(name)) + raise QueryValidationError(f"{name} must be 0 or 1") def _query_list_limit(value, default, maximum): @@ -171,7 +173,7 @@ async def _json_or_form_payload(request): try: return json.loads(body or b"{}"), True except json.JSONDecodeError as e: - raise QueryValidationError("Invalid JSON: {}".format(e)) + raise QueryValidationError(f"Invalid JSON: {e}") return await request.post_vars(), False @@ -192,7 +194,7 @@ async def _analyze_user_query(datasette, db, sql, *, actor): try: analysis = await db.analyze_sql(sql, params) except sqlite3.DatabaseError as ex: - raise QueryValidationError("Could not analyze query: {}".format(ex)) from ex + raise QueryValidationError(f"Could not analyze query: {ex}") from ex is_write = _analysis_is_write(analysis) if is_write: @@ -293,8 +295,7 @@ def _coerce_execute_write_payload(data, is_json): for key, value in data.items(): if key in {"sql", "csrftoken", "_json"}: continue - if key.startswith(SQL_PARAMETER_FORM_PREFIX): - key = key[len(SQL_PARAMETER_FORM_PREFIX) :] + key = key.removeprefix(SQL_PARAMETER_FORM_PREFIX) params[key] = value if not isinstance(params, dict): raise QueryValidationError("params must be a dictionary") @@ -314,7 +315,7 @@ async def _prepare_execute_write(datasette, db, sql, params, actor): try: analysis = await db.analyze_sql(sql, params) except sqlite3.DatabaseError as ex: - raise QueryValidationError("Could not analyze query: {}".format(ex)) from ex + raise QueryValidationError(f"Could not analyze query: {ex}") from ex if not _analysis_is_write(analysis): raise QueryValidationError( "Use /-/query for read-only SQL; this endpoint only executes writes" @@ -496,7 +497,7 @@ async def _inserted_row_url(datasette, db, analysis, cursor): ) try: result = await db.execute( - "select {} from {} where rowid = ?".format(select, escape_sqlite(table)), + f"select {select} from {escape_sqlite(table)} where rowid = ?", [lastrowid], ) except sqlite3.DatabaseError: diff --git a/datasette/views/row.py b/datasette/views/row.py index c90a3bbe..b1388299 100644 --- a/datasette/views/row.py +++ b/datasette/views/row.py @@ -8,34 +8,35 @@ from dataclasses import dataclass, field import markupsafe import sqlite_utils -from datasette.utils.asgi import NotFound, Forbidden, PayloadTooLarge, Response from datasette.database import QueryInterrupted -from datasette.events import UpdateRowEvent, DeleteRowEvent +from datasette.events import DeleteRowEvent, UpdateRowEvent +from datasette.extras import ExtraScope, extra_names_from_request +from datasette.plugins import pm from datasette.resources import TableResource -from .base import BaseView, DatasetteError, stream_csv from datasette.utils import ( + CustomJSONEncoder, + CustomRow, + InvalidSql, + WriteJsonValueError, add_cors_headers, await_me_maybe, call_with_supported_arguments, - CustomJSONEncoder, - CustomRow, decode_write_json_row, - InvalidSql, + escape_sqlite, make_slot_function, path_from_row_pks, path_with_format, path_with_removed_args, - to_css_class, - escape_sqlite, sqlite3, - WriteJsonValueError, + to_css_class, ) -from datasette.plugins import pm -from datasette.extras import extra_names_from_request, ExtraScope +from datasette.utils.asgi import Forbidden, NotFound, PayloadTooLarge, Response + from . import Context, from_extra +from .base import BaseView, DatasetteError, stream_csv from .table import ( - display_columns_and_rows, _table_page_data, + display_columns_and_rows, row_label_from_label_column, ) from .table_extras import RowExtraContext, resolve_row_extras, table_extra_registry @@ -187,16 +188,16 @@ class RowView(BaseView): data, extra_template_data, templates = response_or_template_contexts except QueryInterrupted as ex: raise DatasetteError( - textwrap.dedent(""" + textwrap.dedent(f"""

SQL query took too long. The time limit is controlled by the sql_time_limit_ms configuration option.

- + - """.format(markupsafe.escape(ex.sql))).strip(), + """).strip(), title="SQL Interrupted", status=400, message_is_html=True, @@ -207,15 +208,13 @@ class RowView(BaseView): ) except (sqlite3.OperationalError, InvalidSql) as e: raise DatasetteError(str(e), title="Invalid SQL", status=400) - except sqlite3.OperationalError as e: - raise DatasetteError(str(e)) except DatasetteError: raise end = time.perf_counter() data["query_ms"] = (end - start) * 1000 - if format_ in self.ds.renderers.keys(): + if format_ in self.ds.renderers: # Dispatch request to the correct output format renderer # (CSV is not handled here due to streaming) result = call_with_supported_arguments( @@ -258,7 +257,7 @@ class RowView(BaseView): if status_code is not None: response.status = status_code else: - raise NotFound("Invalid format: {}".format(format_)) + raise NotFound(f"Invalid format: {format_}") ttl = request.args.get("_ttl", None) if ttl is None or not ttl.isdigit(): @@ -373,9 +372,7 @@ class RowView(BaseView): view_name=self.name, ), headers={ - "Link": '<{}>; rel="alternate"; type="application/json+datasette"'.format( - alternate_url_json - ) + "Link": f'<{alternate_url_json}>; rel="alternate"; type="application/json+datasette"' }, ) @@ -500,7 +497,7 @@ class RowView(BaseView): row_action_label = pk_path if row_label and row_label != pk_path: - row_action_label = "{} {}".format(pk_path, row_label) + row_action_label = f"{pk_path} {row_label}" row_action_permissions = {} if is_table and db.is_mutable: @@ -513,7 +510,7 @@ class RowView(BaseView): row_actions = [] if row_action_permissions.get("update-row"): attrs = { - "aria-label": "Edit row {}".format(row_action_label), + "aria-label": f"Edit row {row_action_label}", "data-row": row_path, "data-row-action": "edit", } @@ -529,7 +526,7 @@ class RowView(BaseView): ) if row_action_permissions.get("delete-row"): attrs = { - "aria-label": "Delete row {}".format(row_action_label), + "aria-label": f"Delete row {row_action_label}", "data-row": row_path, "data-row-action": "delete", } @@ -679,7 +676,7 @@ class RowView(BaseView): key, ",".join(pk_values), ) - foreign_key_tables.append({**fk, **{"count": count, "link": link}}) + foreign_key_tables.append({**fk, "count": count, "link": link}) return foreign_key_tables @@ -705,23 +702,21 @@ async def _row_flash_message(db, action, resolved, row=None): if label: label = _truncated_row_flash_label(label) if label and label != pk_label: - return "{} row {} ({})".format(action, pk_label, label) - return "{} row {}".format(action, pk_label) + return f"{action} row {pk_label} ({label})" + return f"{action} row {pk_label}" async def _resolve_row_and_check_permission(datasette, request, permission): - from datasette.app import DatabaseNotFound, TableNotFound, RowNotFound + from datasette.app import DatabaseNotFound, RowNotFound, TableNotFound try: resolved = await datasette.resolve_row(request) except DatabaseNotFound as e: - return False, Response.error( - ["Database not found: {}".format(e.database_name)], 404 - ) + return False, Response.error([f"Database not found: {e.database_name}"], 404) except TableNotFound as e: - return False, Response.error(["Table not found: {}".format(e.table)], 404) + return False, Response.error([f"Table not found: {e.table}"], 404) except RowNotFound as e: - return False, Response.error(["Record not found: {}".format(e.pk_values)], 404) + return False, Response.error([f"Record not found: {e.pk_values}"], 404) # Ensure user has permission to delete this row if not await datasette.allowed( @@ -753,7 +748,8 @@ class RowDeleteView(BaseView): try: await resolved.db.execute_write_fn(delete_row, request=request) - except Exception as e: + except Exception as e: # noqa: BLE001 + # TODO: narrow to expected write errors so Datasette bugs surface as 500s return Response.error([str(e)], 400) await self.ds.track_event( @@ -793,7 +789,7 @@ class RowUpdateView(BaseView): try: data = await request.json() except json.JSONDecodeError as e: - return Response.error(["Invalid JSON: {}".format(e)]) + return Response.error([f"Invalid JSON: {e}"]) except PayloadTooLarge as e: return Response.error([str(e)], 413) @@ -836,7 +832,8 @@ class RowUpdateView(BaseView): try: await resolved.db.execute_write_fn(update_row, request=request) - except Exception as e: + except Exception as e: # noqa: BLE001 + # TODO: narrow to expected write errors so Datasette bugs surface as 500s return Response.error([str(e)], 400) result = {"ok": True} diff --git a/datasette/views/special.py b/datasette/views/special.py index c13191a1..a77f221f 100644 --- a/datasette/views/special.py +++ b/datasette/views/special.py @@ -1,23 +1,25 @@ import json import logging +import secrets +import urllib + +from datasette.events import CreateTokenEvent, LoginEvent, LogoutEvent from datasette.jump import JumpSQL, namespace_sql_params from datasette.plugins import pm -from datasette.events import LogoutEvent, LoginEvent, CreateTokenEvent from datasette.resources import DatabaseResource, TableResource -from datasette.utils.asgi import Response, Forbidden from datasette.utils import ( UNSTABLE_API_MESSAGE, actor_matches_allow, - parse_size_limit, add_cors_headers, await_me_maybe, error_body, - tilde_encode, + parse_size_limit, tilde_decode, + tilde_encode, ) +from datasette.utils.asgi import Forbidden, Response + from .base import BaseView, View -import secrets -import urllib logger = logging.getLogger(__name__) @@ -179,9 +181,7 @@ class AutocompleteDebugView(BaseView): ) context.update( { - "autocomplete_url": "{}/-/autocomplete".format( - self.ds.urls.table(database_name, table_name) - ), + "autocomplete_url": f"{self.ds.urls.table(database_name, table_name)}/-/autocomplete", "label_column": await db.label_column_for_table(table_name), } ) @@ -420,8 +420,11 @@ class AllowedResourcesView(BaseView): row["reason"] = resource.reasons allowed_rows.append(row) - except Exception: - # If catalog tables don't exist yet, return empty results + except Exception: # noqa: BLE001 + # Returns empty results if the catalog tables don't exist yet, but + # also swallows the AttributeError raised for instance-level actions + # such as view-instance, which have no resource_class. + # TODO: handle that case explicitly and narrow this to sqlite3.Error return ( { "ok": True, @@ -523,7 +526,7 @@ class PermissionRulesView(BaseView): from datasette.utils.actions_sql import build_permission_rules_sql - union_sql, union_params, restriction_sqls = await build_permission_rules_sql( + union_sql, union_params, _restriction_sqls = await build_permission_rules_sql( self.ds, actor, action ) await self.ds.refresh_schemas() @@ -600,7 +603,7 @@ class PermissionRulesView(BaseView): async def _check_permission_for_actor(ds, action, parent, child, actor): - """Shared logic for checking permissions. Returns a dict with check results.""" + """Shared logic for checking and explaining a permission decision.""" if action not in ds.actions: return error_body(f"Unknown action: {action}", 404), 404 @@ -629,15 +632,28 @@ async def _check_permission_for_actor(ds, action, parent, child, actor): allowed = await ds.allowed(action=action, resource=resource_obj, actor=actor) + from datasette.utils.actions_sql import explain_permission_for_resource + + explanation = await explain_permission_for_resource( + datasette=ds, + actor=actor, + action=action, + parent=parent, + child=child, + ) + response = { "ok": True, + "unstable": UNSTABLE_API_MESSAGE, "action": action, "allowed": bool(allowed), + "actor": actor, "resource": { "parent": parent, "child": child, "path": _resource_path(parent, child), }, + "explanation": explanation, } if actor and "id" in actor: @@ -655,11 +671,25 @@ class PermissionCheckView(BaseView): as_format = request.url_vars.get("format") if not as_format: + actions = [ + { + "name": action.name, + "description": action.description, + "takes_parent": action.takes_parent, + "takes_child": action.takes_child, + "also_requires": action.also_requires, + } + for action in sorted( + self.ds.actions.values(), key=lambda action: action.name + ) + ] return await self.render( ["debug_check.html"], request, { - "sorted_actions": sorted(self.ds.actions.keys()), + "actions": actions, + "actor_json": request.args.get("actor") + or json.dumps(request.actor, indent=2), "has_debug_permission": True, }, ) @@ -671,9 +701,18 @@ class PermissionCheckView(BaseView): parent = request.args.get("parent") child = request.args.get("child") + actor = request.actor + actor_json = request.args.get("actor") + if actor_json is not None: + try: + actor = json.loads(actor_json) + except json.JSONDecodeError as ex: + return Response.error(f"Invalid actor JSON: {ex}", 400) + if actor is not None and not isinstance(actor, dict): + return Response.error("actor must be a JSON object or null", 400) response, status = await _check_permission_for_actor( - self.ds, action, parent, child, request.actor + self.ds, action, parent, child, actor ) return Response.json(response, status=status) @@ -900,7 +939,7 @@ class ApiExplorerView(BaseView): tables.append({"name": table, "links": table_links}) table_links.append( { - "label": "Get rows for {}".format(table), + "label": f"Get rows for {table}", "method": "GET", "path": self.ds.urls.table(name, table, format="json"), } @@ -920,7 +959,7 @@ class ApiExplorerView(BaseView): { "path": self.ds.urls.table(name, table) + "/-/insert", "method": "POST", - "label": "Insert rows into {}".format(table), + "label": f"Insert rows into {table}", "json": { "rows": [ { @@ -934,7 +973,7 @@ class ApiExplorerView(BaseView): { "path": self.ds.urls.table(name, table) + "/-/upsert", "method": "POST", - "label": "Upsert rows into {}".format(table), + "label": f"Upsert rows into {table}", "json": { "rows": [ { @@ -964,7 +1003,7 @@ class ApiExplorerView(BaseView): table_links.append( { "path": self.ds.urls.table(name, table) + "/-/drop", - "label": "Drop table {}".format(table), + "label": f"Drop table {table}", "json": {"confirm": False}, "method": "POST", } @@ -981,7 +1020,7 @@ class ApiExplorerView(BaseView): database_links.append( { "path": self.ds.urls.database(name) + "/-/create", - "label": "Create table in {}".format(name), + "label": f"Create table in {name}", "json": { "table": "new_table", "columns": [ diff --git a/datasette/views/stored_queries.py b/datasette/views/stored_queries.py index d64f37d2..0bbe9f38 100644 --- a/datasette/views/stored_queries.py +++ b/datasette/views/stored_queries.py @@ -124,7 +124,7 @@ class QueryListView(BaseView): pairs.append(("_next", page.next)) next_url = self.ds.absolute_url( request, - "{}?{}".format(request.path, urlencode(pairs)), + f"{request.path}?{urlencode(pairs)}", ) current_filters = { @@ -415,7 +415,7 @@ class QueryDefinitionView(BaseView): query_name = tilde_decode(request.url_vars["query"]) query = await self.ds.get_query(db.name, query_name) if query is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) if not await self.ds.allowed( action="view-query", resource=QueryResource(db.name, query_name), @@ -439,7 +439,7 @@ class QueryUpdateView(BaseView): query_name = tilde_decode(request.url_vars["query"]) existing = await self.ds.get_query(db.name, query_name) if existing is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) if not await self.ds.allowed( action="update-query", resource=QueryResource(db.name, query_name), @@ -532,7 +532,7 @@ class QueryEditView(BaseView): async def get(self, request): db, query_name, existing = await self._load(request) if existing is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) await self.ds.ensure_permission( action="update-query", resource=QueryResource(db.name, query_name), @@ -545,7 +545,7 @@ class QueryEditView(BaseView): async def post(self, request): db, query_name, existing = await self._load(request) if existing is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) if not await self.ds.allowed( action="update-query", resource=QueryResource(db.name, query_name), @@ -629,7 +629,7 @@ class QueryDeleteView(BaseView): async def get(self, request): db, query_name, existing = await self._load(request) if existing is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) await self.ds.ensure_permission( action="delete-query", resource=QueryResource(db.name, query_name), @@ -653,7 +653,7 @@ class QueryDeleteView(BaseView): async def post(self, request): db, query_name, existing = await self._load(request) if existing is None: - return Response.error(["Query not found: {}".format(query_name)], 404) + return Response.error([f"Query not found: {query_name}"], 404) if not await self.ds.allowed( action="delete-query", resource=QueryResource(db.name, query_name), @@ -665,13 +665,13 @@ class QueryDeleteView(BaseView): ["Trusted queries cannot be deleted using the API"], 403 ) - data, is_json = await _json_or_form_payload(request) + _data, is_json = await _json_or_form_payload(request) await self.ds.remove_query(db.name, query_name) if is_json: return Response.json({"ok": True}) self.ds.add_message( request, - "Query “{}” deleted".format(existing.title or query_name), + f"Query “{existing.title or query_name}” deleted", self.ds.INFO, ) return Response.redirect(self.ds.urls.path(self.ds.urls.database(db.name))) diff --git a/datasette/views/table.py b/datasette/views/table.py index f7edd744..3eb80854 100644 --- a/datasette/views/table.py +++ b/datasette/views/table.py @@ -3,48 +3,51 @@ import itertools import json import urllib import urllib.parse +from dataclasses import dataclass, field import markupsafe +import sqlite_utils +from datasette import tracer from datasette.column_types import SQLiteType -from datasette.extras import extra_names_from_request -from datasette.plugins import pm +from datasette.database import QueryInterrupted from datasette.events import ( AlterTableEvent, DropTableEvent, InsertRowsEvent, UpsertRowsEvent, ) -from datasette.database import QueryInterrupted -from datasette import tracer +from datasette.extras import ExtraScope, extra_names_from_request +from datasette.filters import Filters +from datasette.plugins import pm from datasette.resources import DatabaseResource, TableResource from datasette.utils import ( - add_cors_headers, - await_me_maybe, - call_with_supported_arguments, CustomJSONEncoder, CustomRow, + InvalidSql, + WriteJsonValueError, + add_cors_headers, append_querystring, + await_me_maybe, + call_with_supported_arguments, compound_keys_after_sql, decode_write_json_rows, - format_bytes, - make_slot_function, - tilde_encode, escape_sqlite, filters_should_redirect, + format_bytes, is_url, + make_slot_function, path_from_row_pks, path_with_added_args, path_with_format, path_with_removed_args, path_with_replaced_args, + sqlite3, + tilde_encode, to_css_class, truncate_url, urlsafe_components, value_as_boolean, - InvalidSql, - WriteJsonValueError, - sqlite3, ) from datasette.utils.asgi import ( BadRequest, @@ -54,11 +57,7 @@ from datasette.utils.asgi import ( Request, Response, ) -from datasette.filters import Filters -import sqlite_utils -from dataclasses import dataclass, field -from datasette.extras import ExtraScope from . import Context, from_extra from .base import BaseView, DatasetteError, stream_csv from .database import QueryView @@ -536,7 +535,7 @@ async def _table_insert_ui( columns.append(column_data) data = { - "path": "{}/-/insert".format(datasette.urls.table(database_name, table_name)), + "path": f"{datasette.urls.table(database_name, table_name)}/-/insert", "tableName": table_name, "columns": columns, "bulkColumns": bulk_columns, @@ -544,8 +543,8 @@ async def _table_insert_ui( "maxInsertRows": datasette.setting("max_insert_rows"), } if can_update: - data["upsertPath"] = "{}/-/upsert".format( - datasette.urls.table(database_name, table_name) + data["upsertPath"] = ( + f"{datasette.urls.table(database_name, table_name)}/-/upsert" ) return data @@ -604,7 +603,7 @@ async def _table_alter_ui( columns.append(column_data) data = { - "path": "{}/-/alter".format(datasette.urls.table(database_name, table_name)), + "path": f"{datasette.urls.table(database_name, table_name)}/-/alter", "tableName": table_name, "columns": columns, "primaryKeys": pks, @@ -630,9 +629,7 @@ async def _table_alter_ui( actor=request.actor, ) if can_drop_table: - data["dropPath"] = "{}/-/drop".format( - datasette.urls.table(database_name, table_name) - ) + data["dropPath"] = f"{datasette.urls.table(database_name, table_name)}/-/drop" return data @@ -728,12 +725,10 @@ async def display_columns_and_rows( row_label = row_label_from_label_column(row, label_column) row_action_label = pk_path if row_label and row_label != pk_path: - row_action_label = "{} {}".format(pk_path, row_label) + row_action_label = f"{pk_path} {row_label}" table_path = datasette.urls.table(database_name, table_name) - row_link = '{flat_pks}'.format( - table_path=table_path, - flat_pks=str(markupsafe.escape(pk_path)), - flat_pks_quoted=row_path, + row_link = ( + f'{markupsafe.escape(pk_path)!s}' ) edit_icon = ( '

{}

'.format(error) in response2.text + assert f'

{error}

' in response2.text else: # Check create-token event event = last_event(app_client.ds) @@ -228,7 +231,7 @@ def test_auth_create_token( # And test that token response3 = app_client.get( "/-/actor.json", - headers={"Authorization": "Bearer {}".format("dstok_{}".format(token))}, + headers={"Authorization": "Bearer {}".format(f"dstok_{token}")}, ) assert response3.status == 200 assert response3.json["actor"]["id"] == "test" @@ -241,7 +244,7 @@ async def test_auth_create_token_not_allowed_for_tokens(ds_client): ) response = await ds_client.get( "/-/create-token", - headers={"Authorization": "Bearer dstok_{}".format(ds_tok)}, + headers={"Authorization": f"Bearer dstok_{ds_tok}"}, ) assert response.status_code == 403 @@ -286,12 +289,12 @@ async def test_auth_with_dstok_token(ds_client, scenario, should_work): elif scenario == "invalid_token": token = "invalid" if token: - token = "dstok_{}".format(token) + token = f"dstok_{token}" if scenario == "allow_signed_tokens_off": ds_client.ds._settings["allow_signed_tokens"] = False headers = {} if token: - headers["Authorization"] = "Bearer {}".format(token) + headers["Authorization"] = f"Bearer {token}" response = await ds_client.get("/-/actor.json", headers=headers) try: if should_work: @@ -338,7 +341,7 @@ def test_cli_create_token(app_client, expires): assert details.keys() == expected_keys assert details["a"] == "test" response = app_client.get( - "/-/actor.json", headers={"Authorization": "Bearer {}".format(token)} + "/-/actor.json", headers={"Authorization": f"Bearer {token}"} ) if expires is None or expires > 0: expected_actor = { diff --git a/tests/test_base_view.py b/tests/test_base_view.py index c1b0cf20..b46f7ce1 100644 --- a/tests/test_base_view.py +++ b/tests/test_base_view.py @@ -1,8 +1,10 @@ -from datasette.views.base import View +import json + +import pytest + from datasette import Request, Response from datasette.app import Datasette -import json -import pytest +from datasette.views.base import View class GetView(View): diff --git a/tests/test_cli.py b/tests/test_cli.py index cbd8edad..fbd4a8a9 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -1,23 +1,28 @@ -from .fixtures import ( - make_app_client, - TestClient as _TestClient, - EXPECTED_PLUGINS, -) -from datasette.app import SETTINGS -from datasette.plugins import DEFAULT_PLUGINS, pm -from datasette.cli import cli, serve -from datasette.version import __version__ -from datasette.utils import tilde_encode -from datasette.utils.sqlite import sqlite3 -from click.testing import CliRunner import io import json import pathlib -import pytest import sys import textwrap from unittest import mock +import pytest +from click.testing import CliRunner + +from datasette.app import SETTINGS +from datasette.cli import cli, serve +from datasette.plugins import DEFAULT_PLUGINS, pm +from datasette.utils import tilde_encode +from datasette.utils.sqlite import sqlite3 +from datasette.version import __version__ + +from .fixtures import ( + EXPECTED_PLUGINS, + make_app_client, +) +from .fixtures import ( + TestClient as _TestClient, +) + def test_inspect_cli(app_client): runner = CliRunner() @@ -460,7 +465,7 @@ def test_serve_create(tmpdir): @pytest.mark.parametrize("argument", ("-c", "--config")) @pytest.mark.parametrize("format_", ("json", "yaml")) def test_serve_config(tmpdir, argument, format_): - config_path = tmpdir / "datasette.{}".format(format_) + config_path = tmpdir / f"datasette.{format_}" config_path.write_text( ( "settings:\n default_page_size: 5\n" @@ -513,13 +518,13 @@ def test_weird_database_names(tmpdir, filename): result1 = runner.invoke(cli, [db_path, "--get", "/"]) assert result1.exit_code == 0, result1.output filename_no_stem = filename.rsplit(".", 1)[0] - expected_link = '{}'.format( - tilde_encode(filename_no_stem), filename_no_stem + expected_link = ( + f'{filename_no_stem}' ) assert expected_link in result1.output # Now try hitting that database page result2 = runner.invoke( - cli, [db_path, "--get", "/{}".format(tilde_encode(filename_no_stem))] + cli, [db_path, "--get", f"/{tilde_encode(filename_no_stem)}"] ) assert result2.exit_code == 0, result2.output diff --git a/tests/test_cli_serve_get.py b/tests/test_cli_serve_get.py index fe9416d6..01b84f59 100644 --- a/tests/test_cli_serve_get.py +++ b/tests/test_cli_serve_get.py @@ -1,8 +1,10 @@ +import json +import textwrap + +from click.testing import CliRunner + from datasette.cli import cli from datasette.plugins import pm -from click.testing import CliRunner -import textwrap -import json def test_serve_with_get(tmp_path_factory): @@ -44,9 +46,9 @@ def test_serve_with_get(tmp_path_factory): # Annoyingly that new test plugin stays resident - we need # to manually unregister it to avoid conflict with other tests - to_unregister = [ + to_unregister = next( p for p in pm.get_plugins() if p.__name__ == "init_for_serve_with_get.py" - ][0] + ) pm.unregister(to_unregister) diff --git a/tests/test_cli_serve_server.py b/tests/test_cli_serve_server.py index 47f23c08..b7604bb8 100644 --- a/tests/test_cli_serve_server.py +++ b/tests/test_cli_serve_server.py @@ -1,6 +1,7 @@ +import socket + import httpx import pytest -import socket @pytest.mark.serial diff --git a/tests/test_column_types.py b/tests/test_column_types.py index cd308ec9..50c6daed 100644 --- a/tests/test_column_types.py +++ b/tests/test_column_types.py @@ -1,7 +1,11 @@ import json import logging +import time +import markupsafe +import pytest from bs4 import BeautifulSoup as Soup + from datasette.app import Datasette from datasette.column_types import ( ColumnType, @@ -9,11 +13,7 @@ from datasette.column_types import ( ) from datasette.hookspecs import hookimpl from datasette.plugins import pm -from datasette.utils import error_body, sqlite3 -from datasette.utils import StartupError -import markupsafe -import pytest -import time +from datasette.utils import StartupError, error_body, sqlite3 @pytest.fixture @@ -104,7 +104,7 @@ def write_token(ds, actor_id="root", permissions=None): def _headers(token): return { - "Authorization": "Bearer {}".format(token), + "Authorization": f"Bearer {token}", "Content-Type": "application/json", } diff --git a/tests/test_config_dir.py b/tests/test_config_dir.py index 42c6ae60..00540464 100644 --- a/tests/test_config_dir.py +++ b/tests/test_config_dir.py @@ -1,10 +1,12 @@ import json import pathlib + import pytest from datasette.app import Datasette -from datasette.utils.sqlite import sqlite3 from datasette.utils import StartupError +from datasette.utils.sqlite import sqlite3 + from .fixtures import TestClient as _TestClient PLUGIN = """ @@ -109,7 +111,7 @@ def test_settings(config_dir_client): def test_plugins(config_dir_client): response = config_dir_client.get("/-/plugins.json") assert 200 == response.status - plugins = response.json["plugins"] + plugins = response.json assert "hooray.py" in {p["name"] for p in plugins} assert "non_py_file.txt" not in {p["name"] for p in plugins} assert "mypy_cache" not in {p["name"] for p in plugins} diff --git a/tests/test_crossdb.py b/tests/test_crossdb.py index 11e53224..ffd0870c 100644 --- a/tests/test_crossdb.py +++ b/tests/test_crossdb.py @@ -1,7 +1,9 @@ -from datasette.cli import cli -from click.testing import CliRunner -import urllib import sqlite3 +import urllib + +from click.testing import CliRunner + +from datasette.cli import cli def test_crossdb_join(app_client_two_attached_databases_crossdb_enabled): @@ -40,7 +42,7 @@ def test_crossdb_warning_if_too_many_databases(tmp_path_factory): db_dir = tmp_path_factory.mktemp("dbs") dbs = [] for i in range(11): - path = str(db_dir / "db_{}.db".format(i)) + path = str(db_dir / f"db_{i}.db") conn = sqlite3.connect(path) conn.execute("vacuum") conn.close() diff --git a/tests/test_csrf_middleware.py b/tests/test_csrf_middleware.py index 2fcfb216..6c78f69d 100644 --- a/tests/test_csrf_middleware.py +++ b/tests/test_csrf_middleware.py @@ -44,7 +44,7 @@ async def _run_middleware(scope): await mw(scope, None, send) if inner_called: return ("allowed",) - start = [m for m in sent if m["type"] == "http.response.start"][0] + start = next(m for m in sent if m["type"] == "http.response.start") return ("blocked", start["status"]) diff --git a/tests/test_csv.py b/tests/test_csv.py index a2f03776..7758a3c0 100644 --- a/tests/test_csv.py +++ b/tests/test_csv.py @@ -1,8 +1,10 @@ -from datasette.app import Datasette -from bs4 import BeautifulSoup as Soup -import pytest import urllib.parse +import pytest +from bs4 import BeautifulSoup as Soup + +from datasette.app import Datasette + EXPECTED_TABLE_CSV = """id,content 1,hello 2,world diff --git a/tests/test_custom_pages.py b/tests/test_custom_pages.py index 86cdcc6b..32cfc43d 100644 --- a/tests/test_custom_pages.py +++ b/tests/test_custom_pages.py @@ -1,5 +1,7 @@ import pathlib + import pytest + from .fixtures import make_app_client TEST_TEMPLATE_DIRS = str(pathlib.Path(__file__).parent / "test_templates") diff --git a/tests/test_default_deny.py b/tests/test_default_deny.py index f1e43064..f456a17f 100644 --- a/tests/test_default_deny.py +++ b/tests/test_default_deny.py @@ -1,4 +1,5 @@ import pytest + from datasette.app import Datasette from datasette.resources import DatabaseResource, TableResource diff --git a/tests/test_docs.py b/tests/test_docs.py index 0bcb5e62..16df6a46 100644 --- a/tests/test_docs.py +++ b/tests/test_docs.py @@ -2,20 +2,22 @@ Tests to ensure certain things are documented. """ -from datasette import app, utils +import re +from pathlib import Path + +import pytest + import datasette.fixtures # noqa: F401 +from datasette import app, utils from datasette.app import Datasette from datasette.filters import Filters -from pathlib import Path -import pytest -import re docs_path = Path(__file__).parent.parent / "docs" label_re = re.compile(r"\.\. _([^\s:]+):") def get_headings(content, underline="-"): - heading_re = re.compile(r"(\w+)(\([^)]*\))?\n\{}+\n".format(underline)) + heading_re = re.compile(rf"(\w+)(\([^)]*\))?\n\{underline}+\n") return {h[0] for h in heading_re.findall(content)} diff --git a/tests/test_docs_plugins.py b/tests/test_docs_plugins.py index 613160ac..4a0014b4 100644 --- a/tests/test_docs_plugins.py +++ b/tests/test_docs_plugins.py @@ -1,10 +1,11 @@ # fmt: off # -- start datasette_with_plugin_fixture -- -from datasette import hookimpl -from datasette.app import Datasette import pytest import pytest_asyncio +from datasette import hookimpl +from datasette.app import Datasette + @pytest_asyncio.fixture async def datasette_with_plugin(): diff --git a/tests/test_error_shape.py b/tests/test_error_shape.py index 768814fd..94c9a7c9 100644 --- a/tests/test_error_shape.py +++ b/tests/test_error_shape.py @@ -17,8 +17,10 @@ present and the legacy "title" key must not be. https://github.com/simonw/datasette/issues - 1.0 API consistency """ -import pytest import time + +import pytest + from datasette.app import Datasette from datasette.utils import sqlite3 @@ -86,7 +88,7 @@ async def test_write_api_validation_error_shape(ds_error_shape): "/data/docs/-/insert", json={"rows": [{"nope": 1}, {"also_nope": 2}]}, headers={ - "Authorization": "Bearer {}".format(token), + "Authorization": f"Bearer {token}", "Content-Type": "application/json", }, ) @@ -410,7 +412,7 @@ async def test_expired_token_returns_401(ds_error_shape): ) ) response = await ds_error_shape.client.get( - "/-/actor.json", headers={"Authorization": "Bearer {}".format(token)} + "/-/actor.json", headers={"Authorization": f"Bearer {token}"} ) data = assert_canonical_error(response, 401) assert "expired" in data["error"].lower() @@ -446,7 +448,7 @@ async def test_valid_token_still_authenticates(ds_error_shape): ) ) response = await ds_error_shape.client.get( - "/-/actor.json", headers={"Authorization": "Bearer {}".format(token)} + "/-/actor.json", headers={"Authorization": f"Bearer {token}"} ) assert response.status_code == 200 assert response.json()["actor"]["id"] == "root" @@ -477,7 +479,7 @@ async def test_token_when_signed_tokens_disabled_returns_401(tmp_path_factory): ds.sign({"a": "root", "t": int(time.time())}, namespace="token") ) response = await ds.client.get( - "/-/actor.json", headers={"Authorization": "Bearer {}".format(token)} + "/-/actor.json", headers={"Authorization": f"Bearer {token}"} ) data = assert_canonical_error(response, 401) assert "not enabled" in data["error"] @@ -642,7 +644,7 @@ async def test_query_list_size_rejects_non_integer(ds_client): @pytest.mark.asyncio @pytest.mark.parametrize("endpoint", ("allowed", "rules")) async def test_debug_endpoints_use_size_and_page_parameters(ds_error_shape, endpoint): - base = "/-/{}.json?action=view-instance".format(endpoint) + base = f"/-/{endpoint}.json?action=view-instance" ok = await ds_error_shape.client.get( base + "&_size=1&_page=1", actor={"id": "root"} ) diff --git a/tests/test_extras.py b/tests/test_extras.py index 73b4965e..4e008926 100644 --- a/tests/test_extras.py +++ b/tests/test_extras.py @@ -1,4 +1,5 @@ import asyncio +from typing import ClassVar import pytest @@ -7,7 +8,7 @@ from datasette.extras import Extra, ExtraRegistry, ExtraScope class SlowValueExtra(Extra): description = "Returns context['value'], optionally slowly" - scopes = {ExtraScope.TABLE} + scopes: ClassVar[set[ExtraScope]] = {ExtraScope.TABLE} async def resolve(self, context): if context["slow"]: @@ -17,7 +18,7 @@ class SlowValueExtra(Extra): class DependentExtra(Extra): description = "Depends on slow_value" - scopes = {ExtraScope.TABLE} + scopes: ClassVar[set[ExtraScope]] = {ExtraScope.TABLE} async def resolve(self, context, slow_value): return slow_value + 1 @@ -25,7 +26,7 @@ class DependentExtra(Extra): class InternalOnlyExtra(Extra): description = "Internal extra for HTML templates only" - scopes = {ExtraScope.TABLE} + scopes: ClassVar[set[ExtraScope]] = {ExtraScope.TABLE} public = False async def resolve(self, context): @@ -52,7 +53,7 @@ def _registered_extra_classes(): @pytest.mark.parametrize("cls", _registered_extra_classes(), ids=lambda cls: cls.key()) def test_registered_extras_have_descriptions(cls): # Every registered extra is part of the documented template/JSON contract - assert cls.description, "{} is missing a description".format(cls.__name__) + assert cls.description, f"{cls.__name__} is missing a description" def test_registry_is_built_once_per_scope(): diff --git a/tests/test_facets.py b/tests/test_facets.py index 8c22ffce..b8eb6e61 100644 --- a/tests/test_facets.py +++ b/tests/test_facets.py @@ -1,11 +1,14 @@ +import json + +import pytest + from datasette.app import Datasette from datasette.database import Database -from datasette.facets import Facet, ColumnFacet, ArrayFacet, DateFacet -from datasette.utils.asgi import Request +from datasette.facets import ArrayFacet, ColumnFacet, DateFacet, Facet from datasette.utils import detect_json1 +from datasette.utils.asgi import Request + from .fixtures import make_app_client -import json -import pytest @pytest.mark.asyncio @@ -537,7 +540,7 @@ async def test_facet_size(): for j in range(1, 4): await db.execute_write( "insert into neighbourhoods (city, neighbourhood) values (?, ?)", - ["City {}".format(i), "Neighbourhood {}".format(j)], + [f"City {i}", f"Neighbourhood {j}"], ) response = await ds.client.get( "/test_facet_size/neighbourhoods.json?_extra=suggested_facets" diff --git a/tests/test_filters.py b/tests/test_filters.py index eda9e9a1..8d0f3512 100644 --- a/tests/test_filters.py +++ b/tests/test_filters.py @@ -1,7 +1,8 @@ -from datasette.filters import Filters, through_filters, where_filters, search_filters -from datasette.utils.asgi import Request import pytest +from datasette.filters import Filters, search_filters, through_filters, where_filters +from datasette.utils.asgi import Request + @pytest.mark.parametrize( "args,expected_where,expected_params", diff --git a/tests/test_html.py b/tests/test_html.py index b4c47d80..6a7b4907 100644 --- a/tests/test_html.py +++ b/tests/test_html.py @@ -1,16 +1,19 @@ -from bs4 import BeautifulSoup as Soup -from datasette.app import Datasette -from datasette.utils import allowed_pragmas -from .fixtures import make_app_client -from .utils import assert_footer_links, inner_html import copy import hashlib import json import pathlib -import pytest import re import urllib.parse +import pytest +from bs4 import BeautifulSoup as Soup + +from datasette.app import Datasette +from datasette.utils import allowed_pragmas + +from .fixtures import make_app_client +from .utils import assert_footer_links, inner_html + def test_homepage(app_client_two_attached_databases): response = app_client_two_attached_databases.get("/") @@ -142,9 +145,7 @@ def test_static_mounts_hash_cache_control(): ) incorrect_hash = hashlib.sha256(b"incorrect").hexdigest()[:12] - response = client.get( - "/custom-static/test_html.py?_hash={}".format(incorrect_hash) - ) + response = client.get(f"/custom-static/test_html.py?_hash={incorrect_hash}") assert response.status_code == 200 assert "cache-control" not in response.headers @@ -219,11 +220,9 @@ async def test_disallowed_custom_sql_pragma(ds_client): "/fixtures/-/query?sql=SELECT+*+FROM+pragma_not_on_allow_list('idx52')" ) assert response.status_code == 400 - pragmas = ", ".join("pragma_{}()".format(pragma) for pragma in allowed_pragmas) + pragmas = ", ".join(f"pragma_{pragma}()" for pragma in allowed_pragmas) assert ( - "Statement contained a disallowed PRAGMA. Allowed pragma functions are {}".format( - pragmas - ) + f"Statement contained a disallowed PRAGMA. Allowed pragma functions are {pragmas}" in response.text ) @@ -778,8 +777,8 @@ def test_stored_query_show_hide_metadata_option( }, memory=True, ) as client: - expected_show_hide_fragment = '({})'.format( - expected_show_hide_link, expected_show_hide_text + expected_show_hide_fragment = ( + f'({expected_show_hide_text})' ) response = client.get("/_memory/one" + querystring) html = response.text @@ -788,10 +787,7 @@ def test_stored_query_show_hide_metadata_option( )[0] assert show_hide_fragment == expected_show_hide_fragment if expected_hidden: - assert ( - ''.format(expected_hidden) - in html - ) + assert f'' in html else: assert '; rel="alternate"; type="application/json+datasette"'.format( - expected - ) + assert link == f'<{expected}>; rel="alternate"; type="application/json+datasette"' assert ( - ''.format( - expected - ) + f'' in response.text ) @@ -1292,8 +1284,8 @@ async def test_database_color(ds_client): expected_color = ds_client.ds.get_database("fixtures").color # Should be something like #9403e5 expected_fragments = ( - "10px solid #{}".format(expected_color), - "border-color: #{}".format(expected_color), + f"10px solid #{expected_color}", + f"border-color: #{expected_color}", ) assert len(expected_color) == 6 for path in ( diff --git a/tests/test_internal_db.py b/tests/test_internal_db.py index e1ab51bb..b4bb964d 100644 --- a/tests/test_internal_db.py +++ b/tests/test_internal_db.py @@ -1,6 +1,7 @@ -import pytest import sqlite3 +import pytest + from datasette.utils import escape_sqlite from datasette.utils.internal_db import INTERNAL_DB_SCHEMA_SQL @@ -137,7 +138,7 @@ async def test_internal_foreign_key_references(ds_client): return { row[1] for row in conn.execute( - "PRAGMA table_info({})".format(escape_sqlite(table_name)) + f"PRAGMA table_info({escape_sqlite(table_name)})" ).fetchall() } @@ -147,7 +148,7 @@ async def test_internal_foreign_key_references(ds_client): for _, name in sorted( (row[5], row[1]) for row in conn.execute( - "PRAGMA table_info({})".format(escape_sqlite(table_name)) + f"PRAGMA table_info({escape_sqlite(table_name)})" ).fetchall() if row[5] ) @@ -159,7 +160,7 @@ async def test_internal_foreign_key_references(ds_client): for table_name in table_names: foreign_key_rows = conn.execute( - "PRAGMA foreign_key_list({})".format(escape_sqlite(table_name)) + f"PRAGMA foreign_key_list({escape_sqlite(table_name)})" ).fetchall() foreign_keys_by_id = {} for foreign_key in foreign_key_rows: @@ -169,25 +170,16 @@ async def test_internal_foreign_key_references(ds_client): foreign_key_rows.sort(key=lambda row: row[1]) other_table = foreign_key_rows[0][2] other_columns = [row[4] for row in foreign_key_rows] - message = 'Column "{}.{}" references other table "{}" which does not exist'.format( - table_name, foreign_key_rows[0][3], other_table - ) + message = f'Column "{table_name}.{foreign_key_rows[0][3]}" references other table "{other_table}" which does not exist' assert other_table in table_names, message + " (bad table)" if all(other_column is None for other_column in other_columns): other_columns = primary_keys_for_table(other_table) - length_message = 'Foreign key from "{}" to "{}" has {} columns but references {} columns'.format( - table_name, - other_table, - len(foreign_key_rows), - len(other_columns), - ) + length_message = f'Foreign key from "{table_name}" to "{other_table}" has {len(foreign_key_rows)} columns but references {len(other_columns)} columns' assert len(other_columns) == len(foreign_key_rows), length_message for foreign_key, other_column in zip(foreign_key_rows, other_columns): column = foreign_key[3] - message = 'Column "{}.{}" references other column "{}.{}" which does not exist'.format( - table_name, column, other_table, other_column - ) + message = f'Column "{table_name}.{column}" references other column "{other_table}.{other_column}" which does not exist' assert other_column in columns_by_table[other_table], ( message + " (bad column)" ) @@ -245,10 +237,10 @@ async def test_stale_catalog_entry_database_fix(tmp_path): @pytest.mark.asyncio async def test_stale_catalog_child_entries_removed_for_missing_database(tmp_path): - from datasette.app import Datasette - import sqlite3 + from datasette.app import Datasette + internal_db_path = str(tmp_path / "internal.db") alpha_db_path = str(tmp_path / "alpha.db") bravo_db_path = str(tmp_path / "bravo.db") @@ -293,10 +285,10 @@ async def test_stale_catalog_child_entries_removed_for_missing_database(tmp_path @pytest.mark.asyncio async def test_orphan_stale_catalog_child_entries_removed(tmp_path): - from datasette.app import Datasette - import sqlite3 + from datasette.app import Datasette + internal_db_path = str(tmp_path / "internal.db") alpha_db_path = str(tmp_path / "alpha.db") diff --git a/tests/test_internals_database.py b/tests/test_internals_database.py index bad4e8ca..b1093b1c 100644 --- a/tests/test_internals_database.py +++ b/tests/test_internals_database.py @@ -3,16 +3,23 @@ Tests for the datasette.database.Database class """ import asyncio -from types import SimpleNamespace -from datasette.app import Datasette -from datasette.database import Database, ExecuteWriteResult, Results, MultipleValues -from datasette.database import DatasetteClosedError -from datasette.database import _deliver_write_result -from datasette.utils.sqlite import sqlite3, supports_returning -from datasette.utils import Column -import pytest -import time import uuid +from types import SimpleNamespace + +import pytest +import sqlite_utils + +from datasette.app import Datasette +from datasette.database import ( + Database, + DatasetteClosedError, + ExecuteWriteResult, + MultipleValues, + Results, + _deliver_write_result, +) +from datasette.utils import Column +from datasette.utils.sqlite import sqlite3, supports_returning requires_sqlite_returning = pytest.mark.skipif( not supports_returning(), reason="SQLite does not support RETURNING" @@ -43,7 +50,7 @@ async def test_results_first(db): @pytest.mark.parametrize("expected", (True, False)) async def test_results_bool(db, expected): where = "" if expected else "where pk = 0" - results = await db.execute("select * from facetable {}".format(where)) + results = await db.execute(f"select * from facetable {where}") assert bool(results) is expected @@ -615,7 +622,7 @@ async def test_execute_write_block_false(db): "update roadside_attractions set name = ? where pk = ?", ["Mystery!", 1], ) - time.sleep(0.1) + await asyncio.sleep(0.1) rows = await db.execute("select name from roadside_attractions where pk = 1") assert "Mystery!" == rows.rows[0][0] @@ -633,7 +640,7 @@ async def test_execute_write_with_returning_block_false(db): ) assert isinstance(task_id, uuid.UUID) - time.sleep(0.1) + await asyncio.sleep(0.1) assert ( await db.execute("select name from write_returning_block_false") ).single_value() == "Cleo" @@ -718,15 +725,51 @@ async def test_execute_write_fn_exception(db): await db.execute_write_fn(write_fn) +@pytest.mark.asyncio +@pytest.mark.parametrize("num_sql_threads", (0, 1)) +async def test_execute_write_fn_sqlite_utils_transaction(tmp_path, num_sql_threads): + # A write inside a failing Datasette task must never become visible or + # survive the rollback. Exercise both the synchronous and writer-thread + # paths against a file-backed database so a second connection can observe + # committed state independently. + db_path = tmp_path / "test.db" + sqlite3.connect(db_path).close() + ds = Datasette([str(db_path)], settings={"num_sql_threads": num_sql_threads}) + db = ds.get_database("test") + await db.execute_write("create table items (id integer primary key)") + # This reader is used inside the write callback, which may run on another + # thread, but it is never accessed concurrently. + reader = sqlite3.connect(db_path, check_same_thread=False) + + def insert_then_fail(conn): + # Datasette must open the outer transaction before sqlite-utils writes. + assert conn.in_transaction + sqlite_utils.Database(conn)["items"].insert({"id": 1}) + # If sqlite-utils committed its own transaction, this would return 1. + assert reader.execute("select count(*) from items").fetchone()[0] == 0 + # Simulate a later step failing after the sqlite-utils write succeeded. + raise ValueError("deliberate") + + try: + with pytest.raises(ValueError, match="deliberate"): + await db.execute_write_fn(insert_then_fail) + # The outer transaction must roll back the sqlite-utils write as well. + assert reader.execute("select count(*) from items").fetchone()[0] == 0 + finally: + reader.close() + db.close() + + @pytest.mark.asyncio @pytest.mark.parametrize("param_name", ["conn", "connection", "db", "c"]) async def test_execute_write_fn_accepts_any_single_param_name(db, param_name): # Plugins historically relied on the fact that the callback was invoked # positionally, so any parameter name worked. Preserve that contract. scope = {} - exec( - "def write_fn({0}):\n" - " return {0}.execute('select 1 + 1').fetchone()[0]".format(param_name), + # exec() is how we build a function with a parameterized argument name + exec( # noqa: S102 + f"def write_fn({param_name}):\n" + f" return {param_name}.execute('select 1 + 1').fetchone()[0]", scope, ) write_fn = scope["write_fn"] @@ -750,7 +793,9 @@ async def test_execute_write_fn_with_track_event(db): @pytest.mark.asyncio -@pytest.mark.timeout(1) +# func_only so the budget covers the write-thread call under test, not the +# one-off app_client fixture setup this test may be first to trigger +@pytest.mark.timeout(1, func_only=True) async def test_execute_write_fn_connection_exception(tmpdir, app_client): path = str(tmpdir / "immutable.db") conn = sqlite3.connect(path) diff --git a/tests/test_internals_datasette.py b/tests/test_internals_datasette.py index 85598c05..ed2aeaf0 100644 --- a/tests/test_internals_datasette.py +++ b/tests/test_internals_datasette.py @@ -9,13 +9,15 @@ import importlib import os import sqlite3 import time + +import pytest +from itsdangerous import BadSignature + from datasette import Context -from datasette.app import Datasette, Database, ResourcesSQL +from datasette.app import Database, Datasette, ResourcesSQL from datasette.database import DatasetteClosedError from datasette.resources import DatabaseResource from datasette.utils import PrefixedUrlString -from itsdangerous import BadSignature -import pytest @pytest.fixture @@ -77,9 +79,7 @@ async def test_static_template_function_hashes_core_asset(tmp_path, monkeypatch) template = ds.get_jinja_environment().from_string("{{ static('demo.js') }}") expected_hash = hashlib.sha256(b"const demo = true;").hexdigest()[:12] - assert await template.render_async() == "/-/static/demo.js?_hash={}".format( - expected_hash - ) + assert await template.render_async() == f"/-/static/demo.js?_hash={expected_hash}" assert isinstance(ds.static("demo.js"), PrefixedUrlString) @@ -101,7 +101,7 @@ def test_static_hash_recalculated_when_cache_headers_disabled(tmp_path, monkeypa asset_path.write_bytes(b"let a = 2;") expected_hash = hashlib.sha256(b"let a = 2;").hexdigest()[:12] - assert ds.static("demo.js") == "/-/static/demo.js?_hash={}".format(expected_hash) + assert ds.static("demo.js") == f"/-/static/demo.js?_hash={expected_hash}" assert ds.static("demo.js") != first_url @@ -114,12 +114,12 @@ def test_static_hashes_mounted_static_file(tmp_path): expected_hash = hashlib.sha256(b"body { color: black; }").hexdigest()[:12] assert ds.static("styles.css", mount="assets") == ( - "/assets/styles.css?_hash={}".format(expected_hash) + f"/assets/styles.css?_hash={expected_hash}" ) ds._settings["base_url"] = "/prefix/" assert ds.static("styles.css", mount="assets") == ( - "/prefix/assets/styles.css?_hash={}".format(expected_hash) + f"/prefix/assets/styles.css?_hash={expected_hash}" ) @@ -144,9 +144,7 @@ def test_static_hashes_plugin_static_file(tmp_path, monkeypatch): expected_hash = hashlib.sha256(b"console.log('plugin');").hexdigest()[:12] assert ds.static("plugin.js", plugin="datasette_cluster_map") == ( - "/-/static-plugins/datasette_cluster_map/plugin.js?_hash={}".format( - expected_hash - ) + f"/-/static-plugins/datasette_cluster_map/plugin.js?_hash={expected_hash}" ) diff --git a/tests/test_internals_datasette_client.py b/tests/test_internals_datasette_client.py index e9aaaae8..51b38f8d 100644 --- a/tests/test_internals_datasette_client.py +++ b/tests/test_internals_datasette_client.py @@ -1,6 +1,7 @@ import httpx import pytest import pytest_asyncio + from datasette.app import Datasette @@ -238,7 +239,7 @@ async def test_in_client_returns_false_outside_request(datasette): @pytest.mark.asyncio async def test_in_client_returns_true_inside_request(): """Test that datasette.in_client() returns True inside a client request""" - from datasette import hookimpl, Response + from datasette import Response, hookimpl class TestPlugin: __name__ = "test_in_client_plugin" diff --git a/tests/test_internals_request.py b/tests/test_internals_request.py index 6d2dc70a..e982628b 100644 --- a/tests/test_internals_request.py +++ b/tests/test_internals_request.py @@ -1,7 +1,9 @@ -from datasette.utils.asgi import PayloadTooLarge, Request import json + import pytest +from datasette.utils.asgi import PayloadTooLarge, Request + def _post_scope(headers=None): return { diff --git a/tests/test_internals_response.py b/tests/test_internals_response.py index 2366dcde..aa3e1ae2 100644 --- a/tests/test_internals_response.py +++ b/tests/test_internals_response.py @@ -1,7 +1,9 @@ -from datasette.utils.asgi import Response import json + import pytest +from datasette.utils.asgi import Response + def test_response_html(): response = Response.html("Hello from HTML") diff --git a/tests/test_internals_urls.py b/tests/test_internals_urls.py index 24fa745d..50c61995 100644 --- a/tests/test_internals_urls.py +++ b/tests/test_internals_urls.py @@ -1,6 +1,7 @@ +import pytest + from datasette.app import Datasette from datasette.utils import PrefixedUrlString -import pytest @pytest.fixture(scope="module") diff --git a/tests/test_label_column_for_table.py b/tests/test_label_column_for_table.py index 7667b595..b67b8882 100644 --- a/tests/test_label_column_for_table.py +++ b/tests/test_label_column_for_table.py @@ -1,6 +1,7 @@ import pytest -from datasette.database import Database + from datasette.app import Datasette +from datasette.database import Database @pytest.mark.asyncio diff --git a/tests/test_load_extensions.py b/tests/test_load_extensions.py index cdadb091..61cdb3e0 100644 --- a/tests/test_load_extensions.py +++ b/tests/test_load_extensions.py @@ -1,7 +1,9 @@ -from datasette.app import Datasette -import pytest from pathlib import Path +import pytest + +from datasette.app import Datasette + # not necessarily a full path - the full compiled path looks like "ext.dylib" # or another suffix, but sqlite will, under the hood, decide which file # extension to use based on the operating system (apple=dylib, windows=dll etc) diff --git a/tests/test_messages.py b/tests/test_messages.py index 62d9f647..60eb7938 100644 --- a/tests/test_messages.py +++ b/tests/test_messages.py @@ -1,6 +1,7 @@ -from .utils import cookie_was_deleted import pytest +from .utils import cookie_was_deleted + @pytest.mark.asyncio @pytest.mark.parametrize( diff --git a/tests/test_multipart.py b/tests/test_multipart.py index 0dc3ecd7..ab38bfb7 100644 --- a/tests/test_multipart.py +++ b/tests/test_multipart.py @@ -6,12 +6,12 @@ Uses TDD approach - these tests are written first, then implementation follows. import base64 import json -import pytest from collections import namedtuple +import pytest from multipart_form_data_conformance import get_tests_dir -from datasette.utils.asgi import Request, BadRequest +from datasette.utils.asgi import BadRequest, Request def make_receive(body: bytes): diff --git a/tests/test_package.py b/tests/test_package.py index f05f3ece..43b20589 100644 --- a/tests/test_package.py +++ b/tests/test_package.py @@ -1,9 +1,11 @@ -from click.testing import CliRunner -from datasette import cli -from unittest import mock import os import pathlib +from unittest import mock + import pytest +from click.testing import CliRunner + +from datasette import cli class CaptureDockerfile: diff --git a/tests/test_permission_endpoints.py b/tests/test_permission_endpoints.py index 8726ab62..c54bfbfd 100644 --- a/tests/test_permission_endpoints.py +++ b/tests/test_permission_endpoints.py @@ -6,6 +6,7 @@ Tests for permission endpoints: import pytest import pytest_asyncio + from datasette.app import Datasette @@ -432,8 +433,8 @@ async def test_execute_sql_requires_view_database(): A user who has execute-sql permission but not view-database permission should not be able to execute SQL on that database. """ - from datasette.permissions import PermissionSQL from datasette import hookimpl + from datasette.permissions import PermissionSQL class TestPermissionPlugin: __name__ = "TestPermissionPlugin" diff --git a/tests/test_permissions.py b/tests/test_permissions.py index 88fe577f..73c44682 100644 --- a/tests/test_permissions.py +++ b/tests/test_permissions.py @@ -1,20 +1,23 @@ import collections +import copy +import json +import re +import time +import urllib +from pprint import pprint + +import pytest +import pytest_asyncio from asgiref.sync import async_to_sync +from bs4 import BeautifulSoup as Soup +from click.testing import CliRunner + from datasette.app import Datasette from datasette.cli import cli from datasette.default_permissions import restrictions_allow_action from datasette.utils import UNSTABLE_API_MESSAGE + from .fixtures import assert_permissions_checked, make_app_client -from click.testing import CliRunner -from bs4 import BeautifulSoup as Soup -import copy -import json -from pprint import pprint -import pytest_asyncio -import pytest -import re -import time -import urllib @pytest.fixture(scope="module") @@ -457,6 +460,20 @@ async def test_permissions_debug(ds_client, filter_): assert checks == expected_checks +@pytest.mark.asyncio +@pytest.mark.parametrize( + "permissions_debug,expected_status", + ( + (1, 200), + (0, 403), + ), +) +async def test_permissions_debug_numeric_boolean(permissions_debug, expected_status): + ds = Datasette(config={"permissions": {"permissions-debug": permissions_debug}}) + response = await ds.client.get("/-/permissions") + assert response.status_code == expected_status + + @pytest.mark.asyncio @pytest.mark.parametrize( "actor,allow,expected_fragment", @@ -588,9 +605,7 @@ def test_permissions_cascade(cascade_app_client, path, permissions, expected_sta ) assert ( response.status == expected_status - ), "path: {}, permissions: {}, expected_status: {}, status: {}".format( - path, permissions, expected_status, response.status - ) + ), f"path: {path}, permissions: {permissions}, expected_status: {expected_status}, status: {response.status}" finally: cascade_app_client.ds.config = previous_config @@ -748,7 +763,12 @@ async def test_actor_restricted_permissions( } if actor.get("id"): expected["actor_id"] = actor["id"] - assert response.json() == expected + data = response.json() + for key, value in expected.items(): + assert data[key] == value + assert data["actor"] == actor + assert data["explanation"]["allowed"] is expected_result + assert data["explanation"]["summary"] PermConfigTestCase = collections.namedtuple( @@ -1734,6 +1754,8 @@ async def test_permission_check_view_requires_debug_permission(): data = response.json() assert data["action"] == "view-instance" assert data["allowed"] is True + assert data["explanation"]["allowed"] is True + assert data["explanation"]["summary"] @pytest.mark.asyncio @@ -1759,6 +1781,211 @@ async def test_permission_check_view_query_actions(action): } +@pytest.mark.asyncio +async def test_permission_check_explains_specificity_for_hypothetical_actor(): + ds = Datasette( + config={ + "permissions": {"view-table": {"id": "alice"}}, + "databases": { + "analytics": { + "permissions": {"view-table": False}, + "tables": { + "public": {"permissions": {"view-table": {"id": "alice"}}} + }, + } + }, + } + ) + ds.root_enabled = True + await ds.invoke_startup() + + def path_for(child): + return "/-/check.json?" + urllib.parse.urlencode( + { + "action": "view-table", + "parent": "analytics", + "child": child, + "actor": json.dumps({"id": "alice"}), + } + ) + + public_response = await ds.client.get(path_for("public"), actor={"id": "root"}) + assert public_response.status_code == 200 + public = public_response.json() + assert public["actor"] == {"id": "alice"} + assert public["allowed"] is True + assert public["explanation"]["allowed"] is True + assert public["explanation"]["winning_scope"] == "resource" + public_rules = public["explanation"]["matched_rules"] + assert any( + rule["scope"] == "resource" and rule["effect"] == "allow" and rule["decisive"] + for rule in public_rules + ) + assert any( + rule["scope"] == "parent" + and rule["effect"] == "deny" + and rule["ignored_because"] == "A more specific rule matched" + for rule in public_rules + ) + + private_response = await ds.client.get(path_for("private"), actor={"id": "root"}) + assert private_response.status_code == 200 + private = private_response.json() + assert private["allowed"] is False + assert private["explanation"]["allowed"] is False + assert private["explanation"]["winning_scope"] == "parent" + assert private["explanation"]["summary"].startswith("Denied by a parent-level rule") + + +@pytest.mark.asyncio +async def test_permission_check_explains_deny_wins_at_same_scope(): + ds = Datasette(config={"permissions": {"view-table": {"id": "someone-else"}}}) + ds.root_enabled = True + await ds.invoke_startup() + path = "/-/check.json?" + urllib.parse.urlencode( + { + "action": "view-table", + "parent": "analytics", + "child": "users", + "actor": json.dumps({"id": "alice"}), + } + ) + response = await ds.client.get(path, actor={"id": "root"}) + assert response.status_code == 200 + data = response.json() + assert data["allowed"] is False + assert data["explanation"]["winning_scope"] == "global" + rules = data["explanation"]["matched_rules"] + assert any(rule["effect"] == "deny" and rule["decisive"] for rule in rules) + assert any( + rule["effect"] == "allow" + and rule["ignored_because"] == "A deny rule matched at the same scope" + for rule in rules + ) + + +@pytest.mark.asyncio +async def test_permission_check_explains_default_deny(): + ds = Datasette() + ds.root_enabled = True + await ds.invoke_startup() + path = "/-/check.json?" + urllib.parse.urlencode( + { + "action": "insert-row", + "parent": "analytics", + "child": "users", + "actor": json.dumps({"id": "alice"}), + } + ) + response = await ds.client.get(path, actor={"id": "root"}) + assert response.status_code == 200 + data = response.json() + assert data["allowed"] is False + explanation = data["explanation"] + assert explanation["allowed"] is False + assert explanation["matched_rules"] == [] + assert explanation["winning_scope"] is None + assert explanation["summary"] == ( + "Denied because no permission rule matched this actor and resource." + ) + + +@pytest.mark.asyncio +async def test_permission_check_explains_actor_restrictions(): + ds = Datasette() + ds.root_enabled = True + await ds.invoke_startup() + restricted_actor = { + "id": "alice", + "_r": {"r": {"analytics": {"public": ["vt"]}}}, + } + path = "/-/check.json?" + urllib.parse.urlencode( + { + "action": "view-table", + "parent": "analytics", + "child": "private", + "actor": json.dumps(restricted_actor), + } + ) + response = await ds.client.get(path, actor={"id": "root"}) + assert response.status_code == 200 + data = response.json() + assert data["allowed"] is False + explanation = data["explanation"] + assert explanation["rule_allowed"] is True + assert explanation["restriction_allowed"] is False + assert explanation["allowed"] is False + assert explanation["restrictions"] + assert any( + restriction["allowed"] is False for restriction in explanation["restrictions"] + ) + assert "actor's restrictions" in explanation["summary"] + + +@pytest.mark.asyncio +async def test_permission_check_explains_required_actions(): + from datasette import hookimpl + from datasette.permissions import PermissionSQL + + class StoreQueryPermissions: + @hookimpl + def permission_resources_sql(self, actor, action): + if not actor or actor.get("id") != "alice": + return None + if action == "store-query": + return PermissionSQL( + sql="SELECT 'analytics' AS parent, NULL AS child, 1 AS allow, 'alice can store queries' AS reason" + ) + if action == "execute-sql": + return PermissionSQL( + sql="SELECT 'analytics' AS parent, NULL AS child, 0 AS allow, 'alice cannot execute SQL' AS reason" + ) + + ds = Datasette() + ds.root_enabled = True + await ds.invoke_startup() + ds.pm.register(StoreQueryPermissions(), name="store-query-test") + path = "/-/check.json?" + urllib.parse.urlencode( + { + "action": "store-query", + "parent": "analytics", + "actor": json.dumps({"id": "alice"}), + } + ) + response = await ds.client.get(path, actor={"id": "root"}) + assert response.status_code == 200 + data = response.json() + assert data["allowed"] is False + explanation = data["explanation"] + assert explanation["rule_allowed"] is True + assert explanation["required_actions"][0]["action"] == "execute-sql" + assert explanation["required_actions"][0]["allowed"] is False + assert explanation["summary"] == ( + "Denied because store-query also requires execute-sql, which was denied." + ) + + +@pytest.mark.asyncio +async def test_permission_check_hypothetical_actor_validation(): + ds = Datasette() + ds.root_enabled = True + await ds.invoke_startup() + + response = await ds.client.get( + "/-/check.json?action=view-instance&actor=not-json", + actor={"id": "root"}, + ) + assert response.status_code == 400 + assert response.json()["error"].startswith("Invalid actor JSON:") + + response = await ds.client.get( + "/-/check.json?action=view-instance&actor=%5B%5D", + actor={"id": "root"}, + ) + assert response.status_code == 400 + assert response.json()["error"] == "actor must be a JSON object or null" + + @pytest.mark.asyncio async def test_root_allow_block_with_table_restricted_actor(): """ @@ -1813,7 +2040,7 @@ async def test_databases_json_respects_view_database(tmp_path_factory): paths = [] for name in ("public", "private"): - path = str(db_directory / "{}.db".format(name)) + path = str(db_directory / f"{name}.db") conn = _sqlite3.connect(path) conn.execute("vacuum") conn.close() diff --git a/tests/test_playwright.py b/tests/test_playwright.py index ab9d3568..eb1edb57 100644 --- a/tests/test_playwright.py +++ b/tests/test_playwright.py @@ -410,29 +410,6 @@ def test_datasette_homepage_contains_datasette(page, datasette_server): assert "Datasette" in page.locator("body").inner_text() -@pytest.mark.playwright -def test_table_header_stays_visible_while_scrolling(page, datasette_server): - page.set_viewport_size({"width": 600, "height": 400}) - page.goto(f"{datasette_server}fixtures/sortable?_size=100") - - wrapper = page.locator(".table-wrapper") - header = wrapper.locator("table.rows-and-columns th").first - assert wrapper.evaluate("node => node.scrollWidth > node.clientWidth") - - wrapper.evaluate("""node => window.scrollTo( - 0, - node.getBoundingClientRect().top + window.scrollY + 100 - )""") - page.wait_for_function("""() => Math.abs(document.querySelector( - 'table.rows-and-columns th' - ).getBoundingClientRect().top) < 1""") - assert abs(header.bounding_box()["y"]) < 1 - - wrapper.evaluate("node => node.scrollLeft = 100") - assert wrapper.evaluate("node => node.scrollLeft") == 100 - assert abs(header.bounding_box()["y"]) < 1 - - @pytest.mark.playwright def test_create_table_flow(page, datasette_server): page.goto(f"{datasette_server}data") diff --git a/tests/test_plugins.py b/tests/test_plugins.py index 5c4034db..734f0fc2 100644 --- a/tests/test_plugins.py +++ b/tests/test_plugins.py @@ -1,21 +1,3 @@ -from bs4 import BeautifulSoup as Soup -from .fixtures import ( - make_app_client, - TEMP_PLUGIN_SECRET_FILE, - PLUGINS_DIR, - TestClient as _TestClient, -) # noqa -from click.testing import CliRunner -from datasette.app import Datasette -from datasette import cli, hookimpl -from datasette.fixtures import TABLES -from datasette.filters import FilterArguments -from datasette.plugins import get_plugins, DEFAULT_PLUGINS, pm -from datasette.permissions import PermissionSQL, Action -from datasette.resources import DatabaseResource -from datasette.utils.sqlite import sqlite3 -from datasette.utils import StartupError, await_me_maybe -from jinja2 import ChoiceLoader, FileSystemLoader import base64 import datetime import importlib @@ -24,9 +6,32 @@ import os import pathlib import re import textwrap -import pytest import urllib +import pytest +from bs4 import BeautifulSoup as Soup +from click.testing import CliRunner +from jinja2 import ChoiceLoader, FileSystemLoader + +from datasette import cli, hookimpl +from datasette.app import Datasette +from datasette.filters import FilterArguments +from datasette.fixtures import TABLES +from datasette.permissions import Action, PermissionSQL +from datasette.plugins import DEFAULT_PLUGINS, get_plugins, pm +from datasette.resources import DatabaseResource +from datasette.utils import StartupError, await_me_maybe +from datasette.utils.sqlite import sqlite3 + +from .fixtures import ( + PLUGINS_DIR, + TEMP_PLUGIN_SECRET_FILE, + make_app_client, +) +from .fixtures import ( + TestClient as _TestClient, +) + at_memory_re = re.compile(r" at 0x\w+") @@ -35,7 +40,7 @@ at_memory_re = re.compile(r" at 0x\w+") ) def test_plugin_hooks_have_tests(plugin_hook): """Every plugin hook should be referenced in this test module""" - tests_in_this_module = [t for t in globals().keys() if t.startswith("test_hook_")] + tests_in_this_module = [t for t in globals() if t.startswith("test_hook_")] ok = False for test in tests_in_this_module: if plugin_hook in test: @@ -125,11 +130,11 @@ async def test_hook_extra_css_urls(ds_client, path, expected_decoded_object): response = await ds_client.get(path) assert response.status_code == 200 links = Soup(response.text, "html.parser").find_all("link") - special_href = [ + special_href = next( link for link in links if link.attrs["href"].endswith("/extra-css-urls-demo.css") - ][0]["href"] + )["href"] # This link has a base64-encoded JSON blob in it encoded = special_href.split("/")[3] actual_decoded_object = json.loads(base64.b64decode(encoded).decode("utf8")) @@ -152,7 +157,7 @@ async def test_hook_extra_js_urls(ds_client): "type": "module", }, ]: - assert any(s == attrs for s in script_attrs), "Expected: {}".format(attrs) + assert any(s == attrs for s in script_attrs), f"Expected: {attrs}" @pytest.mark.asyncio @@ -315,7 +320,8 @@ async def test_plugin_config_env_from_list(ds_client): @pytest.mark.asyncio async def test_plugin_config_file(ds_client): - with open(TEMP_PLUGIN_SECRET_FILE, "w") as fp: + # Blocking write is fine here - it is tiny test setup, not request handling + with open(TEMP_PLUGIN_SECRET_FILE, "w") as fp: # noqa: ASYNC230 fp.write("FROM_FILE") assert {"foo": "FROM_FILE"} == ds_client.ds.plugin_config("file-plugin") os.remove(TEMP_PLUGIN_SECRET_FILE) @@ -823,7 +829,7 @@ def test_hook_register_routes_with_datasette(configured_path): assert response.status_code == 200 assert configured_path.upper() == response.text # Other one should 404 - other_path = [p for p in ("path1", "path2") if configured_path != p][0] + other_path = next(p for p in ("path1", "path2") if configured_path != p) assert client.get(f"/{other_path}/", follow_redirects=True).status_code == 404 @@ -928,7 +934,7 @@ async def test_plugin_startup_can_add_queries(): await datasette.add_query( "data", "from_startup", - "select {}".format(result.first()[0]), + f"select {result.first()[0]}", source="plugin", ) @@ -1040,7 +1046,7 @@ async def test_hook_handle_exception(ds_client): @pytest.mark.asyncio @pytest.mark.parametrize("param", ("_custom_error", "_custom_error_async")) async def test_hook_handle_exception_custom_response(ds_client, param): - response = await ds_client.get("/trigger-error?{}=1".format(param)) + response = await ds_client.get(f"/trigger-error?{param}=1") assert response.text == param @@ -1374,7 +1380,7 @@ async def test_hook_register_actions_no_duplicates(duplicate): # This should error: with pytest.raises(StartupError) as ex: await ds.invoke_startup() - assert "Duplicate action {}".format(duplicate) in str(ex.value) + assert f"Duplicate action {duplicate}" in str(ex.value) @pytest.mark.asyncio @@ -1482,7 +1488,7 @@ async def test_plugin_is_installed(): datasette.pm.register(DummyPlugin(), name="DummyPlugin") response = await datasette.client.get("/-/plugins.json") assert response.status_code == 200 - installed_plugins = {p["name"] for p in response.json()["plugins"]} + installed_plugins = {p["name"] for p in response.json()} assert "DummyPlugin" in installed_plugins finally: diff --git a/tests/test_publish_cloudrun.py b/tests/test_publish_cloudrun.py index 6617bc77..aebcaa33 100644 --- a/tests/test_publish_cloudrun.py +++ b/tests/test_publish_cloudrun.py @@ -1,10 +1,12 @@ -from click.testing import CliRunner -from datasette import cli -from unittest import mock import json import os -import pytest import textwrap +from unittest import mock + +import pytest +from click.testing import CliRunner + +from datasette import cli @pytest.mark.serial @@ -70,9 +72,7 @@ def test_publish_cloudrun_prompts_for_service( ), mock.call(f"gcloud builds submit --tag {tag}", shell=True), mock.call( - "gcloud run deploy --allow-unauthenticated --platform=managed --image {} input-service --max-instances 1".format( - tag - ), + f"gcloud run deploy --allow-unauthenticated --platform=managed --image {tag} input-service --max-instances 1", shell=True, ), ] @@ -107,9 +107,7 @@ def test_publish_cloudrun(mock_call, mock_output, mock_which, tmp_path_factory): ), mock.call(f"gcloud builds submit --tag {tag}", shell=True), mock.call( - "gcloud run deploy --allow-unauthenticated --platform=managed --image {} test --max-instances 1".format( - tag - ), + f"gcloud run deploy --allow-unauthenticated --platform=managed --image {tag} test --max-instances 1", shell=True, ), ] @@ -186,13 +184,13 @@ def test_publish_cloudrun_memory_cpu( tag = f"us-docker.pkg.dev/{mock_output.return_value}/datasette/datasette-test" expected_call = ( "gcloud run deploy --allow-unauthenticated --platform=managed" - " --image {} test".format(tag) + f" --image {tag} test" ) expected_build_call = f"gcloud builds submit --tag {tag}" if memory: - expected_call += " --memory {}".format(memory) + expected_call += f" --memory {memory}" if cpu: - expected_call += " --cpu {}".format(cpu) + expected_call += f" --cpu {cpu}" if timeout: expected_build_call += f" --timeout {timeout}" # max_instances defaults to 1 diff --git a/tests/test_publish_heroku.py b/tests/test_publish_heroku.py index cab83654..4302ed94 100644 --- a/tests/test_publish_heroku.py +++ b/tests/test_publish_heroku.py @@ -1,9 +1,11 @@ -from click.testing import CliRunner -from datasette import cli -from unittest import mock import os import pathlib +from unittest import mock + import pytest +from click.testing import CliRunner + +from datasette import cli @pytest.mark.serial diff --git a/tests/test_pytest_autoclose_plugin.py b/tests/test_pytest_autoclose_plugin.py index 3af1aace..9b17d24b 100644 --- a/tests/test_pytest_autoclose_plugin.py +++ b/tests/test_pytest_autoclose_plugin.py @@ -20,6 +20,7 @@ def _run_pytest(tmp_path: Path) -> subprocess.CompletedProcess: cwd=str(tmp_path), capture_output=True, text=True, + check=False, ) diff --git a/tests/test_queries.py b/tests/test_queries.py index ffa948a9..15b7ad0f 100644 --- a/tests/test_queries.py +++ b/tests/test_queries.py @@ -15,37 +15,29 @@ from datasette.utils.sqlite import sqlite3, supports_returning requires_sqlite_returning = pytest.mark.skipif( not supports_returning(), reason="SQLite does not support RETURNING" ) -EXPECTED_CREATE_TABLE_TEMPLATE_SQL = "\n".join( - ( - "create table new_table (", - " id integer primary key,", - " name text", - " -- created text default (datetime('now'))", - ")", - ) -) +EXPECTED_CREATE_TABLE_TEMPLATE_SQL = "create table new_table (\n id integer primary key,\n name text\n -- created text default (datetime('now'))\n)" def _template_option_attributes(html, table): - match = re.search(r'