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 )