diff --git a/adserver/api/serializers.py b/adserver/api/serializers.py index ecfb3f5e..0e01c237 100644 --- a/adserver/api/serializers.py +++ b/adserver/api/serializers.py @@ -86,6 +86,11 @@ class AdDecisionSerializer(serializers.Serializer): user_ip = serializers.CharField(required=False) user_ua = serializers.CharField(required=False) + # Chat/AI prompt text used for embedding-based ad targeting + # When provided, the ad server generates an embedding from this text + # to match against advertiser content for niche targeting + prompt = serializers.CharField(max_length=8000, required=False) + # Used to specify a specific ad or campaign to show (used for debugging mostly) force_ad = serializers.CharField(required=False) # slug force_campaign = serializers.CharField(required=False) # slug diff --git a/adserver/api/views.py b/adserver/api/views.py index bb2abc3e..1b061abf 100644 --- a/adserver/api/views.py +++ b/adserver/api/views.py @@ -340,6 +340,8 @@ def decision(self, request, data): campaign_types=campaign_types, url=url, placement_index=serializer.validated_data.get("placement_index"), + # Prompt text for embedding-based targeting + prompt=serializer.validated_data.get("prompt"), # Debugging parameters ad_slug=serializer.validated_data.get("force_ad"), campaign_slug=serializer.validated_data.get("force_campaign"), diff --git a/adserver/chatdemo/__init__.py b/adserver/chatdemo/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/adserver/chatdemo/embedding.py b/adserver/chatdemo/embedding.py new file mode 100644 index 00000000..5d154a9d --- /dev/null +++ b/adserver/chatdemo/embedding.py @@ -0,0 +1,139 @@ +"""Embedding utilities for generating ad-targeting embeddings from chat prompts.""" + +import hashlib +import logging + +from django.conf import settings +from django.core.cache import cache + + +log = logging.getLogger(__name__) + +# Cache embeddings for 1 hour to avoid redundant API calls +EMBEDDING_CACHE_TIMEOUT = 60 * 60 + +# The OpenAI embedding model to use - text-embedding-3-small is cheap and fast +EMBEDDING_MODEL = "text-embedding-3-small" + + +def get_prompt_embedding(prompt_text): + """ + Generate an embedding vector for the given prompt text using OpenAI. + + Returns a list of floats (the embedding vector) or None on failure. + """ + if not prompt_text or not prompt_text.strip(): + return None + + api_key = getattr(settings, "OPENAI_API_KEY", None) + if not api_key: + log.warning("OPENAI_API_KEY not configured, cannot generate prompt embedding") + return None + + # Check cache first + cache_key = _embedding_cache_key(prompt_text) + cached = cache.get(cache_key) + if cached is not None: + return cached + + try: + import openai + + client = openai.OpenAI(api_key=api_key) + response = client.embeddings.create( + model=EMBEDDING_MODEL, + input=prompt_text.strip()[:8000], # Limit input length + ) + embedding = response.data[0].embedding + + cache.set(cache_key, embedding, EMBEDDING_CACHE_TIMEOUT) + return embedding + + except Exception: + log.exception("Failed to generate embedding for prompt") + return None + + +def cosine_similarity(vec_a, vec_b): + """Compute cosine similarity between two vectors.""" + if not vec_a or not vec_b or len(vec_a) != len(vec_b): + return 0.0 + + dot_product = sum(a * b for a, b in zip(vec_a, vec_b)) + magnitude_a = sum(a * a for a in vec_a) ** 0.5 + magnitude_b = sum(b * b for b in vec_b) ** 0.5 + + if magnitude_a == 0 or magnitude_b == 0: + return 0.0 + + return dot_product / (magnitude_a * magnitude_b) + + +def get_prompt_niche_weights(prompt_text, flights): + """ + Compute niche targeting weights for flights based on prompt embedding similarity. + + This mirrors the ethicalads_ext.embedding.utils.get_niche_weights interface + but works directly from prompt text instead of URL content. + + Returns a dict mapping Advertiser -> distance (lower = more similar). + """ + prompt_embedding = get_prompt_embedding(prompt_text) + if not prompt_embedding: + return {} + + weights = {} + + for flight in flights: + advertiser = flight.campaign.advertiser + + # Skip if we already computed for this advertiser + if advertiser in weights: + continue + + # Get the advertiser's embedding if stored, or generate from ad text + ad_embedding = _get_advertiser_embedding(flight) + if not ad_embedding: + continue + + similarity = cosine_similarity(prompt_embedding, ad_embedding) + # Convert similarity to distance (lower = better match) + # niche_targeting threshold compares distance < goal + distance = 1.0 - similarity + weights[advertiser] = distance + + return weights + + +def _get_advertiser_embedding(flight): + """ + Get or generate an embedding for a flight's advertiser content. + + Uses the flight's advertisement text as the content to embed. + """ + # Build a text representation from the flight's ads + ad_texts = [] + for ad in flight.advertisements.filter(live=True)[:5]: + parts = [] + if ad.headline: + parts.append(ad.headline) + if ad.content: + parts.append(ad.content) + if ad.cta: + parts.append(ad.cta) + if not parts and ad.text: + parts.append(ad.text) + if parts: + ad_texts.append(" ".join(parts)) + + if not ad_texts: + return None + + combined_text = " | ".join(ad_texts) + return get_prompt_embedding(combined_text) + + +def _embedding_cache_key(text): + """Generate a cache key for an embedding.""" + text_hash = hashlib.md5(text.strip().lower().encode()).hexdigest() + return f"prompt-embedding-{text_hash}" diff --git a/adserver/chatdemo/urls.py b/adserver/chatdemo/urls.py new file mode 100644 index 00000000..8b6fffd5 --- /dev/null +++ b/adserver/chatdemo/urls.py @@ -0,0 +1,14 @@ +"""URL configuration for the chat demo.""" + +from django.urls import path + +from .views import ChatCompletionProxyView +from .views import ChatDemoView + + +app_name = "chatdemo" + +urlpatterns = [ + path("", ChatDemoView.as_view(), name="chat-demo"), + path("completion/", ChatCompletionProxyView.as_view(), name="chat-completion"), +] diff --git a/adserver/chatdemo/views.py b/adserver/chatdemo/views.py new file mode 100644 index 00000000..6c75d2d1 --- /dev/null +++ b/adserver/chatdemo/views.py @@ -0,0 +1,88 @@ +"""Views for the AI chat demo with ethical ad targeting.""" + +import json +import logging + +from django.conf import settings +from django.http import JsonResponse +from django.utils.decorators import method_decorator +from django.views import View +from django.views.decorators.csrf import csrf_exempt +from django.views.generic import TemplateView + + +log = logging.getLogger(__name__) + + +class ChatDemoView(TemplateView): + """Serve the chat demo HTML page.""" + + template_name = "adserver/chatdemo/chat.html" + + def get_context_data(self, **kwargs): + context = super().get_context_data(**kwargs) + context["publisher_slug"] = getattr( + settings, "ADSERVER_CHAT_DEMO_PUBLISHER", "" + ) + return context + + +@method_decorator(csrf_exempt, name="dispatch") +class ChatCompletionProxyView(View): + """ + Proxy chat completion requests to OpenAI. + + This keeps the OpenAI API key on the server side + and uses a cheap model (gpt-4o-mini) for completions. + """ + + OPENAI_MODEL = "gpt-4o-mini" + + def post(self, request): + api_key = getattr(settings, "OPENAI_API_KEY", None) + if not api_key: + return JsonResponse( + {"error": "OpenAI API key not configured on the server"}, + status=500, + ) + + try: + body = json.loads(request.body) + except (json.JSONDecodeError, ValueError): + return JsonResponse({"error": "Invalid JSON"}, status=400) + + messages = body.get("messages", []) + if not messages: + return JsonResponse({"error": "No messages provided"}, status=400) + + # Limit conversation length to prevent abuse + messages = messages[:50] + + try: + import openai + + client = openai.OpenAI(api_key=api_key) + response = client.chat.completions.create( + model=self.OPENAI_MODEL, + messages=messages, + max_tokens=1024, + temperature=0.7, + ) + + return JsonResponse( + { + "content": response.choices[0].message.content, + "model": response.model, + "usage": { + "prompt_tokens": response.usage.prompt_tokens, + "completion_tokens": response.usage.completion_tokens, + }, + } + ) + + except Exception: + log.exception("OpenAI chat completion failed") + return JsonResponse( + {"error": "Chat completion request failed"}, + status=502, + ) diff --git a/adserver/decisionengine/backends.py b/adserver/decisionengine/backends.py index 148bddf6..f2b5dc85 100644 --- a/adserver/decisionengine/backends.py +++ b/adserver/decisionengine/backends.py @@ -119,6 +119,9 @@ def __init__(self, request, placements, publisher, **kwargs): self.ad_slug = kwargs.get("ad_slug") self.campaign_slug = kwargs.get("campaign_slug") + # Chat/AI prompt text for embedding-based targeting + self.prompt = kwargs.get("prompt") or "" + self.niche_weights = None def get_analyzer_keywords(self): @@ -420,19 +423,29 @@ def select_flight(self): # Apply niche targeting only when any flight has it. # This is to track whether we should do expensive distance queries. - if ( - flights_with_niche_targeting - and "ethicalads_ext.embedding" in settings.INSTALLED_APPS - ): - # We have to do this here, - # so we can filter by the weight in the filter_flight call below - from ethicalads_ext.embedding.utils import get_niche_weights # noqa + if flights_with_niche_targeting: + if self.prompt: + # Use prompt-based embedding for niche targeting + # This enables AI chat contexts to target ads + from ..chatdemo.embedding import get_prompt_niche_weights + + self.niche_weights = get_prompt_niche_weights( + self.prompt, flights_with_niche_targeting + ) + if self.niche_weights: + log.debug( + "Prompt niche targeting weights: %s", + self.niche_weights, + ) + elif "ethicalads_ext.embedding" in settings.INSTALLED_APPS: + # Fall back to URL-based niche targeting + from ethicalads_ext.embedding.utils import get_niche_weights # noqa - self.niche_weights = get_niche_weights( - url=self.url, flights=flights_with_niche_targeting - ) - if self.niche_weights: - log.debug("Niche targeting weights: %s", self.niche_weights) + self.niche_weights = get_niche_weights( + url=self.url, flights=flights_with_niche_targeting + ) + if self.niche_weights: + log.debug("Niche targeting weights: %s", self.niche_weights) for flight in possible_flights: # Handle excluding flights based on targeting diff --git a/adserver/tasks.py b/adserver/tasks.py index 8324a6fe..d842933d 100644 --- a/adserver/tasks.py +++ b/adserver/tasks.py @@ -1484,3 +1484,35 @@ def run_publisher_importers(): """ # PSF is the only importer for now.. psf.run_import(sync=True) + + +@app.task() +def publish_celery_queue_depth(): + """Publish Celery queue depth to CloudWatch for autoscaling the celery ASG.""" + import boto3 + import redis + + broker_url = settings.CELERY_BROKER_URL + queues = ["celery", "analyzer", "priority"] + + try: + r = redis.Redis.from_url(broker_url, socket_timeout=5, socket_connect_timeout=5) + total = sum(r.llen(q) for q in queues) + except Exception: + log.exception("Failed to read Celery queue depths from Redis") + return + + try: + client = boto3.client("cloudwatch", region_name=settings.AWS_S3_REGION_NAME) + client.put_metric_data( + Namespace="EthicalAds/Celery", + MetricData=[ + { + "MetricName": "QueueDepth", + "Value": total, + "Unit": "Count", + } + ], + ) + except Exception: + log.exception("Failed to publish queue depth to CloudWatch") diff --git a/adserver/templates/adserver/chatdemo/chat.html b/adserver/templates/adserver/chatdemo/chat.html new file mode 100644 index 00000000..5011a648 --- /dev/null +++ b/adserver/templates/adserver/chatdemo/chat.html @@ -0,0 +1,520 @@ + + + + + + AI Chat Demo - Ethical Ad Targeting + + + + + +
+
+

AI Chat with Ethical Ad Targeting

+
Chat prompts are used to contextually target relevant ads via embeddings
+
+
+ +
+ + +
+ +
+
+
+
+ Welcome! Start chatting and relevant ads will appear in the sidebar, + targeted based on the context of your conversation. +
+
+ +
Thinking...
+ +
+ + +
+
+ +
+
Contextual Ads
+
+
+ Ads will appear here once you start chatting. + They are targeted based on the content of your conversation. +
+
+
+ Ads are contextually targeted using prompt embeddings — no personal data is collected. +
+
+
+ + + + + diff --git a/adserver/tests/test_chatdemo.py b/adserver/tests/test_chatdemo.py new file mode 100644 index 00000000..bfabc70a --- /dev/null +++ b/adserver/tests/test_chatdemo.py @@ -0,0 +1,307 @@ +"""Tests for the chatdemo app: embedding utilities, views, and related task.""" + +import json +import sys +from unittest import mock + +from django.test import Client +from django.test import TestCase +from django.test import override_settings +from django.urls import reverse + +from ..chatdemo.embedding import _embedding_cache_key +from ..chatdemo.embedding import cosine_similarity +from ..chatdemo.embedding import get_prompt_embedding +from ..chatdemo.embedding import get_prompt_niche_weights +from .common import BaseAdModelsTestCase + + +def _ensure_mock_module(name): + """Inject a MagicMock into sys.modules so ``import `` succeeds.""" + if name not in sys.modules: + sys.modules[name] = mock.MagicMock() + + +# Ensure optional third-party packages are importable even if not installed. +_ensure_mock_module("openai") +_ensure_mock_module("boto3") +_ensure_mock_module("redis") + + +class CosineSimTest(TestCase): + """Tests for the cosine_similarity helper.""" + + def test_identical_vectors(self): + vec = [1.0, 0.0, 0.0] + self.assertAlmostEqual(cosine_similarity(vec, vec), 1.0) + + def test_orthogonal_vectors(self): + self.assertAlmostEqual( + cosine_similarity([1, 0, 0], [0, 1, 0]), + 0.0, + ) + + def test_opposite_vectors(self): + self.assertAlmostEqual( + cosine_similarity([1, 0], [-1, 0]), + -1.0, + ) + + def test_empty_or_mismatched(self): + self.assertEqual(cosine_similarity([], []), 0.0) + self.assertEqual(cosine_similarity(None, None), 0.0) + self.assertEqual(cosine_similarity([1], [1, 2]), 0.0) + + def test_zero_magnitude(self): + self.assertEqual(cosine_similarity([0, 0], [1, 1]), 0.0) + + +class EmbeddingCacheKeyTest(TestCase): + """Tests for _embedding_cache_key.""" + + def test_deterministic(self): + key1 = _embedding_cache_key("hello world") + key2 = _embedding_cache_key("hello world") + self.assertEqual(key1, key2) + + def test_case_insensitive(self): + self.assertEqual( + _embedding_cache_key("Hello"), + _embedding_cache_key("hello"), + ) + + def test_strips_whitespace(self): + self.assertEqual( + _embedding_cache_key(" hello "), + _embedding_cache_key("hello"), + ) + + +class GetPromptEmbeddingTest(TestCase): + """Tests for get_prompt_embedding.""" + + def test_empty_input_returns_none(self): + self.assertIsNone(get_prompt_embedding("")) + self.assertIsNone(get_prompt_embedding(None)) + self.assertIsNone(get_prompt_embedding(" ")) + + @override_settings(OPENAI_API_KEY=None) + def test_no_api_key_returns_none(self): + self.assertIsNone(get_prompt_embedding("some text")) + + @override_settings(OPENAI_API_KEY="test-key") + @mock.patch("adserver.chatdemo.embedding.cache") + def test_cache_hit(self, mock_cache): + mock_cache.get.return_value = [0.1, 0.2, 0.3] + result = get_prompt_embedding("cached text") + self.assertEqual(result, [0.1, 0.2, 0.3]) + + @override_settings(OPENAI_API_KEY="test-key") + @mock.patch("adserver.chatdemo.embedding.cache") + def test_openai_call(self, mock_cache): + mock_cache.get.return_value = None # cache miss + + mock_embedding_data = mock.MagicMock() + mock_embedding_data.embedding = [0.5, 0.6, 0.7] + mock_response = mock.MagicMock() + mock_response.data = [mock_embedding_data] + + with mock.patch("openai.OpenAI") as mock_openai_cls: + mock_client = mock.MagicMock() + mock_openai_cls.return_value = mock_client + mock_client.embeddings.create.return_value = mock_response + + result = get_prompt_embedding("test prompt") + + self.assertEqual(result, [0.5, 0.6, 0.7]) + mock_cache.set.assert_called_once() + + @override_settings(OPENAI_API_KEY="test-key") + @mock.patch("adserver.chatdemo.embedding.cache") + def test_openai_exception_returns_none(self, mock_cache): + mock_cache.get.return_value = None + + with mock.patch("openai.OpenAI") as mock_openai_cls: + mock_client = mock.MagicMock() + mock_openai_cls.return_value = mock_client + mock_client.embeddings.create.side_effect = Exception("API down") + + result = get_prompt_embedding("fail text") + + self.assertIsNone(result) + + +class GetPromptNicheWeightsTest(BaseAdModelsTestCase): + """Tests for get_prompt_niche_weights.""" + + @mock.patch("adserver.chatdemo.embedding.get_prompt_embedding") + def test_no_embedding_returns_empty(self, mock_embed): + mock_embed.return_value = None + result = get_prompt_niche_weights("test", [self.flight]) + self.assertEqual(result, {}) + + @mock.patch("adserver.chatdemo.embedding._get_advertiser_embedding") + @mock.patch("adserver.chatdemo.embedding.get_prompt_embedding") + def test_returns_distance_dict(self, mock_embed, mock_ad_embed): + mock_embed.return_value = [1.0, 0.0, 0.0] + mock_ad_embed.return_value = [1.0, 0.0, 0.0] + + result = get_prompt_niche_weights("test", [self.flight]) + self.assertIn(self.advertiser, result) + self.assertAlmostEqual(result[self.advertiser], 0.0) # identical = distance 0 + + @mock.patch("adserver.chatdemo.embedding._get_advertiser_embedding") + @mock.patch("adserver.chatdemo.embedding.get_prompt_embedding") + def test_no_ad_embedding_skips(self, mock_embed, mock_ad_embed): + mock_embed.return_value = [1.0, 0.0] + mock_ad_embed.return_value = None + + result = get_prompt_niche_weights("test", [self.flight]) + self.assertEqual(result, {}) + + +class ChatDemoViewTest(TestCase): + """Tests for the ChatDemoView template view.""" + + @override_settings(ADSERVER_CHAT_DEMO_PUBLISHER="test-pub") + def test_get_renders(self): + url = reverse("chatdemo:chat-demo") + resp = self.client.get(url) + self.assertEqual(resp.status_code, 200) + self.assertContains(resp, "test-pub") + + +class ChatCompletionProxyViewTest(TestCase): + """Tests for the ChatCompletionProxyView.""" + + def setUp(self): + self.url = reverse("chatdemo:chat-completion") + self.client = Client() + + @override_settings(OPENAI_API_KEY=None) + def test_no_api_key(self): + resp = self.client.post( + self.url, + data=json.dumps({"messages": [{"role": "user", "content": "hi"}]}), + content_type="application/json", + ) + self.assertEqual(resp.status_code, 500) + self.assertIn("error", resp.json()) + + @override_settings(OPENAI_API_KEY="test-key") + def test_invalid_json(self): + resp = self.client.post( + self.url, + data="not json", + content_type="application/json", + ) + self.assertEqual(resp.status_code, 400) + + @override_settings(OPENAI_API_KEY="test-key") + def test_no_messages(self): + resp = self.client.post( + self.url, + data=json.dumps({"messages": []}), + content_type="application/json", + ) + self.assertEqual(resp.status_code, 400) + + @override_settings(OPENAI_API_KEY="test-key") + def test_successful_completion(self): + mock_choice = mock.MagicMock() + mock_choice.message.content = "Hello from AI" + mock_usage = mock.MagicMock() + mock_usage.prompt_tokens = 10 + mock_usage.completion_tokens = 5 + mock_response = mock.MagicMock() + mock_response.choices = [mock_choice] + mock_response.model = "gpt-4o-mini" + mock_response.usage = mock_usage + + with mock.patch("openai.OpenAI") as mock_openai_cls: + mock_client = mock.MagicMock() + mock_openai_cls.return_value = mock_client + mock_client.chat.completions.create.return_value = mock_response + + resp = self.client.post( + self.url, + data=json.dumps({"messages": [{"role": "user", "content": "hi"}]}), + content_type="application/json", + ) + + self.assertEqual(resp.status_code, 200) + data = resp.json() + self.assertEqual(data["content"], "Hello from AI") + self.assertEqual(data["model"], "gpt-4o-mini") + + @override_settings(OPENAI_API_KEY="test-key") + def test_openai_failure(self): + with mock.patch("openai.OpenAI") as mock_openai_cls: + mock_client = mock.MagicMock() + mock_openai_cls.return_value = mock_client + mock_client.chat.completions.create.side_effect = Exception("fail") + + resp = self.client.post( + self.url, + data=json.dumps({"messages": [{"role": "user", "content": "hi"}]}), + content_type="application/json", + ) + + self.assertEqual(resp.status_code, 502) + + +class PublishCeleryQueueDepthTest(TestCase): + """Tests for the publish_celery_queue_depth task.""" + + @mock.patch("boto3.client") + @mock.patch("redis.Redis.from_url") + @override_settings( + CELERY_BROKER_URL="redis://localhost:6379/0", + AWS_S3_REGION_NAME="us-east-1", + ) + def test_publishes_metric(self, mock_redis_from_url, mock_boto_client): + from ..tasks import publish_celery_queue_depth + + mock_redis = mock.MagicMock() + mock_redis.llen.return_value = 5 + mock_redis_from_url.return_value = mock_redis + + mock_cw = mock.MagicMock() + mock_boto_client.return_value = mock_cw + + publish_celery_queue_depth() + + mock_cw.put_metric_data.assert_called_once() + call_kwargs = mock_cw.put_metric_data.call_args[1] + self.assertEqual(call_kwargs["Namespace"], "EthicalAds/Celery") + self.assertEqual(call_kwargs["MetricData"][0]["Value"], 15) # 5 * 3 queues + + @mock.patch("redis.Redis.from_url") + @override_settings(CELERY_BROKER_URL="redis://localhost:6379/0") + def test_redis_failure(self, mock_redis_from_url): + from ..tasks import publish_celery_queue_depth + + mock_redis_from_url.side_effect = Exception("connection failed") + + # Should not raise + publish_celery_queue_depth() + + @mock.patch("boto3.client") + @mock.patch("redis.Redis.from_url") + @override_settings( + CELERY_BROKER_URL="redis://localhost:6379/0", + AWS_S3_REGION_NAME="us-east-1", + ) + def test_cloudwatch_failure(self, mock_redis_from_url, mock_boto_client): + from ..tasks import publish_celery_queue_depth + + mock_redis = mock.MagicMock() + mock_redis.llen.return_value = 0 + mock_redis_from_url.return_value = mock_redis + + mock_cw = mock.MagicMock() + mock_cw.put_metric_data.side_effect = Exception("cloudwatch fail") + mock_boto_client.return_value = mock_cw + + # Should not raise + publish_celery_queue_depth() diff --git a/config/settings/base.py b/config/settings/base.py index 3f19335d..6ac72732 100644 --- a/config/settings/base.py +++ b/config/settings/base.py @@ -598,6 +598,12 @@ ADSERVER_HTTPS = False # Should be True in most production setups ADSERVER_STICKY_DECISION_DURATION = 0 +# OpenAI API key for chat demo completions and embedding generation +OPENAI_API_KEY = env("OPENAI_API_KEY", default=None) + +# The default publisher slug for the chat demo +ADSERVER_CHAT_DEMO_PUBLISHER = env("ADSERVER_CHAT_DEMO_PUBLISHER", default="") + # For customer support emails ADSERVER_SUPPORT_TO_EMAIL = env("ADSERVER_SUPPORT_TO_EMAIL", default=None) ADSERVER_SUPPORT_FORM_ACTION = env("ADSERVER_SUPPORT_FORM_ACTION", default=None) diff --git a/config/settings/production.py b/config/settings/production.py index 78ea127f..f73955ca 100644 --- a/config/settings/production.py +++ b/config/settings/production.py @@ -269,6 +269,11 @@ "task": "adserver.tasks.run_publisher_importers", "schedule": crontab(hour="1", minute="0"), }, + # Publish queue depth to CloudWatch for celery ASG autoscaling + "every-minute-publish-queue-depth": { + "task": "adserver.tasks.publish_celery_queue_depth", + "schedule": crontab(), # Every minute + }, } # Tasks which should only be run if the analyzer is installed diff --git a/config/urls.py b/config/urls.py index 6619c7cf..764e8115 100644 --- a/config/urls.py +++ b/config/urls.py @@ -62,6 +62,8 @@ ] urlpatterns += [ + # AI Chat demo with ethical ad targeting + path(r"chat/", include("adserver.chatdemo.urls")), # Allauth overrides # Disable managing emails for now path( diff --git a/scripts/prep_chat_demo_db.py b/scripts/prep_chat_demo_db.py new file mode 100644 index 00000000..0fce5209 --- /dev/null +++ b/scripts/prep_chat_demo_db.py @@ -0,0 +1,226 @@ +""" +Prep local database for the AI chat demo. + +Usage: + python manage.py shell < scripts/prep_chat_demo_db.py + +This script uses the first existing object for each model where possible, +creating only what's missing. It wires everything together so that the +chat demo at /chat/ can successfully request and display ads. +""" + +import datetime + +from adserver.constants import HOUSE_CAMPAIGN +from adserver.models import AdType +from adserver.models import Advertisement +from adserver.models import Advertiser +from adserver.models import Campaign +from adserver.models import Flight +from adserver.models import Publisher +from adserver.models import PublisherGroup + + +def run(): + print("=" * 60) + print("Prepping database for AI chat demo") + print("=" * 60) + + # --- Publisher Group --- + pub_group = PublisherGroup.objects.first() + if not pub_group: + pub_group = PublisherGroup.objects.create( + name="Chat Demo Group", + slug="chat-demo-group", + default_enabled=True, + ) + print(f" Created PublisherGroup: {pub_group}") + else: + print(f" Using existing PublisherGroup: {pub_group} (slug={pub_group.slug})") + + # --- Publisher --- + publisher = Publisher.objects.first() + if not publisher: + publisher = Publisher.objects.create( + name="Chat Demo Publisher", + slug="chat-demo", + ) + print(f" Created Publisher: {publisher}") + else: + print(f" Using existing Publisher: {publisher} (slug={publisher.slug})") + + # Ensure publisher is configured for the demo + changed_fields = [] + if not publisher.unauthed_ad_decisions: + publisher.unauthed_ad_decisions = True + changed_fields.append("unauthed_ad_decisions") + if not publisher.allow_paid_campaigns: + publisher.allow_paid_campaigns = True + changed_fields.append("allow_paid_campaigns") + if not publisher.allow_house_campaigns: + publisher.allow_house_campaigns = True + changed_fields.append("allow_house_campaigns") + if not publisher.allow_community_campaigns: + publisher.allow_community_campaigns = True + changed_fields.append("allow_community_campaigns") + if not publisher.allow_api_keywords: + publisher.allow_api_keywords = True + changed_fields.append("allow_api_keywords") + if publisher.disabled: + publisher.disabled = False + changed_fields.append("disabled") + if changed_fields: + publisher.save(update_fields=changed_fields) + print(f" Updated Publisher fields: {changed_fields}") + + # Link publisher to publisher group + if not publisher.publisher_groups.filter(pk=pub_group.pk).exists(): + pub_group.publishers.add(publisher) + print(" Added Publisher to PublisherGroup") + + # --- AdType --- + ad_type = AdType.objects.filter(slug="readthedocs-sidebar").first() + if not ad_type: + ad_type = AdType.objects.first() + if not ad_type: + ad_type = AdType.objects.create( + name="Sidebar Image", + slug="image-v1", + has_image=True, + image_width=240, + image_height=180, + has_text=True, + max_text_length=150, + default_enabled=True, + ) + print(f" Created AdType: {ad_type}") + else: + print(f" Using existing AdType: {ad_type} (slug={ad_type.slug})") + + # --- Advertiser --- + advertiser = Advertiser.objects.first() + if not advertiser: + advertiser = Advertiser.objects.create( + name="Demo Advertiser", + slug="demo-advertiser", + ) + print(f" Created Advertiser: {advertiser}") + else: + print(f" Using existing Advertiser: {advertiser} (slug={advertiser.slug})") + + # --- Campaign --- + campaign = Campaign.objects.first() + if not campaign: + campaign = Campaign.objects.create( + name="Demo Campaign", + slug="demo-campaign", + advertiser=advertiser, + campaign_type=HOUSE_CAMPAIGN, + ) + print(f" Created Campaign: {campaign}") + else: + print(f" Using existing Campaign: {campaign} (slug={campaign.slug})") + + # Ensure the campaign's publisher groups include ours + if not campaign.publisher_groups.filter(pk=pub_group.pk).exists(): + campaign.publisher_groups.add(pub_group) + print(" Added PublisherGroup to Campaign") + + # --- Flight --- + flight = Flight.objects.first() + if not flight: + flight = Flight.objects.create( + name="Demo Flight", + slug="demo-flight", + campaign=campaign, + live=True, + cpc=0, + cpm=0, + sold_clicks=10000, + start_date=datetime.date.today() - datetime.timedelta(days=1), + end_date=datetime.date.today() + datetime.timedelta(days=365), + ) + print(f" Created Flight: {flight}") + else: + print(f" Using existing Flight: {flight} (slug={flight.slug})") + + # Ensure flight is live and has a valid date range + changed_fields = [] + if not flight.live: + flight.live = True + changed_fields.append("live") + if flight.start_date > datetime.date.today(): + flight.start_date = datetime.date.today() - datetime.timedelta(days=1) + changed_fields.append("start_date") + if flight.end_date < datetime.date.today(): + flight.end_date = datetime.date.today() + datetime.timedelta(days=365) + changed_fields.append("end_date") + if flight.sold_clicks == 0 and flight.sold_impressions == 0: + flight.sold_clicks = 10000 + changed_fields.append("sold_clicks") + if changed_fields: + flight.save(update_fields=changed_fields) + print(f" Updated Flight fields: {changed_fields}") + + # --- Advertisement --- + ad = Advertisement.objects.first() + if not ad: + ad = Advertisement.objects.create( + name="Demo Ad", + slug="demo-ad", + flight=flight, + headline="Try Our Developer Tools", + content="Build faster with our modern dev platform.", + cta="Learn More", + link="https://example.com/?utm_source=ethicalads", + live=True, + ) + print(f" Created Advertisement: {ad}") + else: + print(f" Using existing Advertisement: {ad} (slug={ad.slug})") + + # Ensure ad is live + if not ad.live: + ad.live = True + ad.save(update_fields=["live"]) + print(" Set Advertisement live=True") + + # Ensure ad has the ad type + if not ad.ad_types.filter(pk=ad_type.pk).exists(): + ad.ad_types.add(ad_type) + print(f" Added AdType '{ad_type.slug}' to Advertisement") + + # --- Summary --- + print() + print("=" * 60) + print("Setup complete! Here's your configuration:") + print("=" * 60) + print() + print(f" Publisher slug: {publisher.slug}") + print(f" AdType slug: {ad_type.slug}") + print(f" Advertiser: {advertiser.name}") + print(f" Campaign: {campaign.name} (type={campaign.campaign_type})") + print(f" Flight: {flight.name} (live={flight.live})") + print(f" Advertisement: {ad.name} (live={ad.live})") + print() + print("Environment variables to set:") + print(f' export ADSERVER_CHAT_DEMO_PUBLISHER="{publisher.slug}"') + print(' export OPENAI_API_KEY="sk-your-key-here"') + print() + print("Then run:") + print(" python manage.py runserver") + print() + print("And open: http://localhost:8000/chat/") + print() + print("Test the ad API directly:") + print( + f" curl 'http://localhost:8000/api/v1/decision/" + f"?publisher={publisher.slug}" + f"&div_ids=ad1" + f"&ad_types={ad_type.slug}" + f"&keywords=python'" + ) + print() + + +run()