summaryrefslogtreecommitdiff
path: root/packages/rss/lektor_rss.py
diff options
context:
space:
mode:
Diffstat (limited to 'packages/rss/lektor_rss.py')
-rw-r--r--packages/rss/lektor_rss.py260
1 files changed, 260 insertions, 0 deletions
diff --git a/packages/rss/lektor_rss.py b/packages/rss/lektor_rss.py
new file mode 100644
index 0000000..224d6f8
--- /dev/null
+++ b/packages/rss/lektor_rss.py
@@ -0,0 +1,260 @@
+from __future__ import annotations
+
+from datetime import date, datetime, time, timezone
+from typing import Any
+
+from feedgenerator.django.utils.feedgenerator import Rss201rev2Feed
+from lektor.build_programs import BuildProgram
+from lektor.context import get_ctx, url_to
+from lektor.db import F
+from lektor.environment import Expression
+from lektor.pluginsystem import Plugin
+from lektor.sourceobj import VirtualSourceObject
+from lektor.utils import build_url
+
+
+def get_field(record: Any, field: str, default: Any = None) -> Any:
+ if field in record:
+ return record[field]
+ return default
+
+
+def get_item_title(item: Any, field: str) -> str:
+ if field in item:
+ return str(item[field])
+ return item.record_label
+
+
+def get_item_body(item: Any, field: str) -> str:
+ if field not in item:
+ raise RuntimeError(f"Body field {field!r} not found in {item!r}")
+
+ with get_ctx().changed_base_url(item.url_path):
+ body = item[field]
+
+ if hasattr(body, "__html__"):
+ return str(body.__html__())
+
+ return str(body)
+
+
+def get_item_date(item: Any, field: str) -> datetime:
+ value = get_field(item, field)
+
+ if value is None:
+ raise RuntimeError(f"Publication date field {field!r} not found in {item!r}")
+
+ if isinstance(value, datetime):
+ if value.tzinfo is None:
+ return value.replace(tzinfo=timezone.utc)
+ return value
+
+ if isinstance(value, date):
+ return datetime.combine(value, time.min, tzinfo=timezone.utc)
+
+ raise TypeError(f"Publication date field {field!r} must be a date or datetime")
+
+
+class RssFeedSource(VirtualSourceObject):
+ def __init__(
+ self,
+ parent: Any,
+ feed_id: str,
+ plugin: RssPlugin,
+ ) -> None:
+ super().__init__(parent)
+ self.feed_id = feed_id
+ self.plugin = plugin
+
+ @property
+ def path(self) -> str:
+ return f"{self.parent.path}@rss/{self.feed_id}"
+
+ @property
+ def url_path(self) -> str:
+ configured_path = self.plugin.get_rss_config(
+ self.feed_id,
+ "url_path",
+ )
+
+ if configured_path:
+ return configured_path
+
+ return build_url([self.parent.url_path, self.filename])
+
+ @property
+ def feed_name(self) -> str:
+ return self.plugin.get_rss_config(self.feed_id, "name") or self.feed_id
+
+ def iter_source_filenames(self):
+ return self.parent.iter_source_filenames()
+
+ def __getattr__(self, name: str) -> Any:
+ try:
+ return self.plugin.get_rss_config(self.feed_id, name)
+ except KeyError as exc:
+ raise AttributeError(name) from exc
+
+
+class RssFeedBuilderProgram(BuildProgram):
+ def produce_artifacts(self) -> None:
+ self.declare_artifact(
+ self.source.url_path,
+ sources=list(self.source.iter_source_filenames()),
+ )
+
+ def build_artifact(self, artifact: Any) -> None:
+ ctx = get_ctx()
+ source = self.source
+ blog = source.parent
+
+ summary = get_field(
+ blog,
+ source.blog_summary_field,
+ "",
+ )
+
+ if hasattr(summary, "__html__"):
+ summary = str(summary.__html__())
+ else:
+ summary = str(summary)
+
+ blog_author = str(
+ get_field(
+ blog,
+ source.blog_author_field,
+ "",
+ )
+ or ""
+ )
+
+ feed = Rss201rev2Feed(
+ title=source.feed_name,
+ link=url_to(blog, external=True),
+ description=summary,
+ feed_url=url_to(source, external=True),
+ language="en",
+ )
+
+ if source.items:
+ expression = Expression(ctx.env, source.items)
+ items = expression.evaluate(ctx.pad)
+ else:
+ items = blog.children
+
+ if source.item_model:
+ items = items.filter(F._model == source.item_model)
+
+ items = items.order_by(f"-{source.item_date_field}").limit(int(source.limit))
+
+ for item in items:
+ item_url = url_to(item, external=True)
+
+ item_author = (
+ get_field(
+ item,
+ source.item_author_field,
+ )
+ or blog_author
+ )
+
+ feed.add_item(
+ title=get_item_title(
+ item,
+ source.item_title_field,
+ ),
+ link=item_url,
+ description=get_item_body(
+ item,
+ source.item_body_field,
+ ),
+ author_name=str(item_author),
+ pubdate=get_item_date(
+ item,
+ source.item_date_field,
+ ),
+ unique_id=item_url,
+ unique_id_is_permalink=True,
+ )
+
+ with artifact.open("wb") as file:
+ feed.write(file, "utf-8")
+
+
+class RssPlugin(Plugin):
+ name = "RSS"
+ description = "Generate RSS 2.0 feeds."
+
+ defaults = {
+ "source_path": "/",
+ "name": None,
+ "url_path": None,
+ "filename": "feed.rss",
+ "blog_author_field": "author",
+ "blog_summary_field": "summary",
+ "items": None,
+ "limit": 50,
+ "item_title_field": "title",
+ "item_body_field": "body",
+ "item_author_field": "author",
+ "item_date_field": "pub_date",
+ "item_model": None,
+ }
+
+ def get_rss_config(
+ self,
+ feed_id: str,
+ key: str,
+ ) -> Any:
+ return self.get_config().get(
+ f"{feed_id}.{key}",
+ self.defaults[key],
+ )
+
+ def on_setup_env(self, **extra: Any) -> None:
+ self.env.add_build_program(
+ RssFeedSource,
+ RssFeedBuilderProgram,
+ )
+
+ @self.env.virtualpathresolver("rss")
+ def rss_path_resolver(
+ node: Any,
+ pieces: list[str],
+ ) -> RssFeedSource | None:
+ if len(pieces) != 1:
+ return None
+
+ feed_id = pieces[0]
+
+ if feed_id not in self.get_config().sections():
+ return None
+
+ source_path = self.get_rss_config(
+ feed_id,
+ "source_path",
+ )
+
+ if node.path != source_path:
+ return None
+
+ return RssFeedSource(
+ node,
+ feed_id,
+ self,
+ )
+
+ @self.env.generator
+ def generate_rss_feeds(source: Any):
+ for feed_id in self.get_config().sections():
+ source_path = self.get_rss_config(
+ feed_id,
+ "source_path",
+ )
+
+ if source.path == source_path:
+ yield RssFeedSource(
+ source,
+ feed_id,
+ self,
+ )