v0.165.0
  1from typing import Any
  2
  3import jinja2
  4from jinja2 import nodes
  5from jinja2.ext import Extension
  6from jinja2.nodes import CallBlock, Node
  7from jinja2.parser import Parser
  8from jinja2.runtime import Context
  9from plain.runtime import settings
 10from plain.templates import register_template_extension
 11from plain.templates.jinja.extensions import InclusionTagExtension
 12
 13
 14@register_template_extension
 15class HTMXJSExtension(InclusionTagExtension):
 16    tags = {"htmx_js"}  # noqa: RUF012 — jinja2 types `tags` as an instance attribute; ClassVar here fails ty's LSP check
 17    template_name = "htmx/js.html"
 18
 19    def get_context(
 20        self, context: Context, *args: Any, **kwargs: Any
 21    ) -> dict[str, Any]:
 22        request = context.get("request")
 23        return {
 24            "DEBUG": settings.DEBUG,
 25            "extensions": kwargs.get("extensions", []),
 26            "csp_nonce": request.csp_nonce if request else None,
 27        }
 28
 29
 30class _FragmentFound(Exception):
 31    """Raised to short-circuit template rendering once the target fragment is found."""
 32
 33    def __init__(self, content: str) -> None:
 34        self.content = content
 35
 36
 37@register_template_extension
 38class HTMXFragmentExtension(Extension):
 39    tags = {"htmxfragment"}  # noqa: RUF012 — jinja2 types `tags` as an instance attribute; ClassVar here fails ty's LSP check
 40
 41    def parse(self, parser: Parser) -> Node:
 42        lineno = next(parser.stream).lineno
 43
 44        fragment_name = parser.parse_expression()
 45
 46        kwargs = []
 47
 48        while parser.stream.current.type != "block_end":
 49            if parser.stream.current.type == "name":
 50                key = parser.stream.current.value
 51                parser.stream.skip()
 52                parser.stream.expect("assign")
 53                value = parser.parse_expression()
 54                kwargs.append(nodes.Keyword(key, value))
 55
 56        body = parser.parse_statements(("name:endhtmxfragment",), drop_needle=True)
 57
 58        call = self.call_method(
 59            "_render_htmx_fragment",
 60            args=[fragment_name, nodes.ContextReference()],
 61            kwargs=kwargs,
 62        )
 63
 64        callblock = CallBlock(call, [], [], body)
 65        callblock.set_lineno(lineno)
 66
 67        return callblock
 68
 69    def _render_htmx_fragment(
 70        self, fragment_name: str, context: dict[str, Any], caller: Any, **kwargs: Any
 71    ) -> str:
 72        # Two-phase fragment targeting (see render_template_fragment):
 73        # Phase 1 skips non-target bodies, phase 2 renders them for nesting.
 74        # Once the target is found, "found" is set so child fragments render
 75        # normally with their wrapper divs.
 76        target_state = context.get("_htmx_target_fragment")
 77        if target_state is not None and not target_state["found"]:
 78            if str(fragment_name) == target_state["name"]:
 79                target_state["found"] = True
 80                content = caller()
 81                raise _FragmentFound(content)
 82            elif target_state["render_bodies"]:
 83                return caller()
 84            else:
 85                return ""
 86
 87        def attrs_to_str(attrs: dict[str, Any]) -> str:
 88            parts = []
 89            for k, v in attrs.items():
 90                if v == "":
 91                    parts.append(k)
 92                else:
 93                    parts.append(f'{k}="{v}"')
 94            return " ".join(parts)
 95
 96        render_lazy = kwargs.get("lazy", False)
 97        as_element = kwargs.get("as", "div")
 98        attrs = {}
 99        for k, v in kwargs.items():
100            if k in ("lazy", "as"):
101                continue
102            if k.startswith("hx_"):
103                attrs[k.replace("_", "-")] = v
104            else:
105                attrs[k] = v
106
107        if render_lazy:
108            attrs.setdefault("hx-trigger", "load from:body")
109            attrs.setdefault("hx-swap", "outerHTML")
110            attrs.setdefault("hx-target", "this")
111            attrs.setdefault("hx-indicator", "this")
112            attrs_str = attrs_to_str(attrs)
113            return f'<{as_element} plain-hx-fragment="{fragment_name}" hx-get {attrs_str}></{as_element}>'
114        else:
115            # Swap innerHTML so we can re-run hx calls inside the fragment automatically
116            attrs.setdefault("hx-swap", "innerHTML")
117            attrs.setdefault("hx-target", "this")
118            attrs.setdefault("hx-indicator", "this")
119            # Add an id that you can use to target the fragment from outside the fragment
120            attrs.setdefault("id", f"plain-hx-fragment-{fragment_name}")
121            attrs_str = attrs_to_str(attrs)
122            return f'<{as_element} plain-hx-fragment="{fragment_name}" {attrs_str}>{caller()}</{as_element}>'
123
124
125def render_template_fragment(
126    *, template: jinja2.Template, fragment_name: str, context: dict[str, Any]
127) -> str:
128    """Render only the named fragment from a template.
129
130    Two-phase approach:
131    1. Skip non-target fragment bodies (fast — handles top-level and loop fragments)
132    2. If not found, render bodies too (handles fragments nested inside other fragments)
133
134    Raises _FragmentFound to short-circuit as soon as the target is found.
135    """
136    for render_bodies in (False, True):
137        target_state = {
138            "name": fragment_name,
139            "found": False,
140            "render_bodies": render_bodies,
141        }
142        try:
143            template.render({**context, "_htmx_target_fragment": target_state})
144        except _FragmentFound as e:
145            return e.content
146
147    raise jinja2.TemplateNotFound(
148        f"Fragment '{fragment_name}' not found in template {template.name}"
149    )