Skip to content

Commit 25334e8

Browse files
committed
Fix Discord feedback message selection
1 parent 75873ca commit 25334e8

2 files changed

Lines changed: 202 additions & 11 deletions

File tree

tools/feedback/discord-feedback

Lines changed: 61 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,10 @@ PWN_COLLEGE_GITHUB_ORG = "pwncollege"
4141
TEXT_CHANNEL_TYPES = {0, 5}
4242
PUBLIC_THREAD_TYPES = {10, 11}
4343
THREAD_PARENT_TYPES = {0, 5, 15, 16}
44+
# Keep ordinary posts, replies, application-command responses, thread starters, and context
45+
# menu responses. Other Discord message types are generated system events rather than
46+
# learner-authored discussion.
47+
FEEDBACK_MESSAGE_TYPES = {0, 19, 20, 21, 23}
4448
MAX_DISCORD_PAGE_SIZE = 100
4549
MAX_DISCORD_SERVER_ERROR_RETRIES = 5
4650
# Image attachments are downloaded so the analysis agent can view screenshots that carry
@@ -640,11 +644,10 @@ def fetch_messages_since(
640644
api: DiscordAPI,
641645
channel: dict[str, Any],
642646
cutoff: datetime.datetime,
643-
max_messages: int,
644647
) -> list[dict[str, Any]]:
645648
messages: list[dict[str, Any]] = []
646649
params: dict[str, Any] = {"limit": MAX_DISCORD_PAGE_SIZE}
647-
while len(messages) < max_messages:
650+
while True:
648651
try:
649652
page = api.request(f"/channels/{channel['id']}/messages", params)
650653
except (DiscordPermissionError, DiscordServerError) as error:
@@ -667,13 +670,51 @@ def fetch_messages_since(
667670
oldest_time = (
668671
message_time if oldest_time is None else min(oldest_time, message_time)
669672
)
670-
if message_time >= cutoff:
673+
if message_time >= cutoff and message_has_feedback_content(message):
671674
messages.append(message)
672675

673676
if oldest_id is None or oldest_time is None or oldest_time < cutoff:
674-
return messages[:max_messages]
677+
return messages
675678
params["before"] = str(oldest_id)
676-
return messages[:max_messages]
679+
680+
681+
def message_has_feedback_content(message: dict[str, Any]) -> bool:
682+
"""Whether a Discord message carries content useful to the feedback analysis.
683+
684+
Discord system events such as member joins have an author and timestamp but no actual
685+
message payload. Do not let those empty events crowd learner feedback out of the global
686+
retention limit.
687+
"""
688+
if message.get("type", 0) not in FEEDBACK_MESSAGE_TYPES:
689+
return False
690+
return bool(
691+
(message.get("content") or "").strip()
692+
or message.get("attachments")
693+
or message.get("embeds")
694+
)
695+
696+
697+
def retain_most_recent_messages(
698+
messages_by_channel: list[tuple[dict[str, Any], list[dict[str, Any]]]],
699+
max_messages: int,
700+
) -> list[tuple[dict[str, Any], list[dict[str, Any]]]]:
701+
"""Apply the message limit globally after every channel has been scanned."""
702+
candidates = [
703+
(channel, message)
704+
for channel, messages in messages_by_channel
705+
for message in messages
706+
]
707+
candidates.sort(
708+
key=lambda item: (parse_iso8601(item[1]["timestamp"]), int(item[1]["id"])),
709+
reverse=True,
710+
)
711+
712+
retained_by_channel: OrderedDict[
713+
str, tuple[dict[str, Any], list[dict[str, Any]]]
714+
] = OrderedDict()
715+
for channel, message in candidates[:max_messages]:
716+
retained_by_channel.setdefault(channel["id"], (channel, []))[1].append(message)
717+
return list(retained_by_channel.values())
677718

678719

679720
def sanitize_content(content: str, channels_by_id: dict[str, dict[str, Any]]) -> str:
@@ -4432,14 +4473,10 @@ def feedback_command(
44324473
channels, channels_by_id = discover_channels(
44334474
api, guild_id, selected_channel_ids, since, include_threads
44344475
)
4435-
click.echo(f"Discovered {len(channels)} readable channel(s)/thread(s)")
4476+
click.echo(f"Discovered {len(channels)} channel(s)/thread(s) to scan")
44364477
messages_by_channel: list[tuple[dict[str, Any], list[dict[str, Any]]]] = []
4437-
remaining = max_messages
44384478
for channel in channels:
4439-
if remaining <= 0:
4440-
break
4441-
messages = fetch_messages_since(api, channel, since, remaining)
4442-
remaining -= len(messages)
4479+
messages = fetch_messages_since(api, channel, since)
44434480
if messages:
44444481
messages_by_channel.append((channel, messages))
44454482
click.echo(
@@ -4453,6 +4490,17 @@ def feedback_command(
44534490
channel_display_name(channel, channels_by_id),
44544491
)
44554492

4493+
matching_message_count = sum(
4494+
len(channel_messages) for _channel, channel_messages in messages_by_channel
4495+
)
4496+
messages_by_channel = retain_most_recent_messages(
4497+
messages_by_channel, max_messages
4498+
)
4499+
click.echo(
4500+
f"Scanned all {len(channels)} channel(s)/thread(s); retaining the most "
4501+
f"recent {min(matching_message_count, max_messages)} of "
4502+
f"{matching_message_count} non-empty message(s)"
4503+
)
44564504
messages = sanitize_messages(
44574505
messages_by_channel,
44584506
channels_by_id,
@@ -4513,6 +4561,8 @@ def feedback_command(
45134561
"since": since.isoformat(),
45144562
"until": until.isoformat(),
45154563
"messages": len(messages),
4564+
"matching_messages": matching_message_count,
4565+
"max_messages": max_messages,
45164566
"image_attachments": image_count,
45174567
"channels": len(channels),
45184568
"include_threads": include_threads,

tools/feedback/tests/test_discord_feedback.py

Lines changed: 141 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -210,6 +210,147 @@ def test_completed_scrape_is_checkpointed_before_fetch_only_returns(self):
210210
self.assertTrue((artifact_dir / "run.json").is_file())
211211

212212

213+
class DiscordScrapeTests(unittest.TestCase):
214+
@staticmethod
215+
def message(message_id, timestamp, content="", **overrides):
216+
message = {
217+
"id": str(message_id),
218+
"timestamp": timestamp,
219+
"content": content,
220+
"attachments": [],
221+
"embeds": [],
222+
"author": {"id": f"author-{message_id}"},
223+
}
224+
message.update(overrides)
225+
return message
226+
227+
def test_fetch_scans_to_cutoff_and_discards_empty_system_events(self):
228+
cutoff = datetime.datetime(2026, 7, 24, tzinfo=datetime.timezone.utc)
229+
channel = {"id": "123", "name": "help"}
230+
api = mock.Mock()
231+
api.request.side_effect = [
232+
[
233+
self.message(
234+
5,
235+
"2026-07-24T00:05:00+00:00",
236+
type=7,
237+
),
238+
self.message(4, "2026-07-24T00:04:00+00:00", "new feedback"),
239+
],
240+
[
241+
self.message(3, "2026-07-24T00:03:00+00:00", "more feedback"),
242+
self.message(2, "2026-07-23T23:59:00+00:00", "before cutoff"),
243+
],
244+
]
245+
246+
messages = discord_feedback.fetch_messages_since(api, channel, cutoff)
247+
248+
self.assertEqual([message["id"] for message in messages], ["4", "3"])
249+
self.assertEqual(api.request.call_count, 2)
250+
self.assertEqual(api.request.call_args_list[1].args[1]["before"], "4")
251+
252+
def test_feedback_content_includes_text_attachments_and_embeds(self):
253+
timestamp = "2026-07-24T00:00:00+00:00"
254+
255+
self.assertFalse(
256+
discord_feedback.message_has_feedback_content(
257+
self.message(1, timestamp, " ", type=7)
258+
)
259+
)
260+
self.assertFalse(
261+
discord_feedback.message_has_feedback_content(
262+
self.message(5, timestamp, "rendered system text", type=6)
263+
)
264+
)
265+
self.assertTrue(
266+
discord_feedback.message_has_feedback_content(
267+
self.message(2, timestamp, "learner feedback")
268+
)
269+
)
270+
self.assertTrue(
271+
discord_feedback.message_has_feedback_content(
272+
self.message(3, timestamp, attachments=[{"id": "attachment"}])
273+
)
274+
)
275+
self.assertTrue(
276+
discord_feedback.message_has_feedback_content(
277+
self.message(4, timestamp, embeds=[{"title": "feedback"}])
278+
)
279+
)
280+
281+
def test_cli_scans_every_channel_then_keeps_global_newest_messages(self):
282+
with tempfile.TemporaryDirectory() as temporary_directory:
283+
repo = pathlib.Path(temporary_directory)
284+
now = datetime.datetime(
285+
2026, 7, 24, 1, 0, 0, tzinfo=datetime.timezone.utc
286+
)
287+
run_id = now.strftime("%Y%m%d-%H%M%S")
288+
artifact_dir = repo / ".discord-feedback" / run_id
289+
channels = [
290+
{"id": "channel-a", "name": "a"},
291+
{"id": "channel-b", "name": "b"},
292+
{"id": "channel-c", "name": "c"},
293+
]
294+
messages = {
295+
"channel-a": [self.message(1, "2026-07-24T00:10:00+00:00", "oldest")],
296+
"channel-b": [self.message(4, "2026-07-24T00:40:00+00:00", "newest")],
297+
"channel-c": [
298+
self.message(2, "2026-07-24T00:20:00+00:00", "older"),
299+
self.message(3, "2026-07-24T00:30:00+00:00", "newer"),
300+
],
301+
}
302+
303+
with (
304+
mock.patch.object(discord_feedback, "git_root", return_value=repo),
305+
mock.patch.object(discord_feedback, "utc_now", return_value=now),
306+
mock.patch.object(
307+
discord_feedback,
308+
"resolve_scrape_since",
309+
return_value=now - datetime.timedelta(hours=1),
310+
),
311+
mock.patch.object(
312+
discord_feedback,
313+
"discover_channels",
314+
return_value=(channels, {item["id"]: item for item in channels}),
315+
),
316+
mock.patch.object(
317+
discord_feedback,
318+
"fetch_messages_since",
319+
side_effect=lambda _api, channel, _since: messages[channel["id"]],
320+
) as fetch_messages_since,
321+
mock.patch.dict(os.environ, {"DISCORD_BOT_TOKEN": "token"}),
322+
):
323+
result = CliRunner().invoke(
324+
discord_feedback.feedback_command,
325+
[
326+
"--fetch-only",
327+
"--no-pr-feedback",
328+
"--max-messages",
329+
"2",
330+
],
331+
)
332+
333+
self.assertEqual(result.exit_code, 0, result.output)
334+
self.assertEqual(fetch_messages_since.call_count, len(channels))
335+
self.assertEqual(
336+
[call.args[1]["id"] for call in fetch_messages_since.call_args_list],
337+
["channel-a", "channel-b", "channel-c"],
338+
)
339+
retained = [
340+
json.loads(line)
341+
for line in (artifact_dir / "messages.jsonl").read_text().splitlines()
342+
]
343+
self.assertEqual([message["id"] for message in retained], ["3", "4"])
344+
run = json.loads((artifact_dir / "run.json").read_text())
345+
self.assertEqual(run["matching_messages"], 4)
346+
self.assertEqual(run["messages"], 2)
347+
self.assertEqual(run["max_messages"], 2)
348+
self.assertIn(
349+
"Scanned all 3 channel(s)/thread(s); retaining the most recent 2 of 4",
350+
result.output,
351+
)
352+
353+
213354
class ValidationTests(unittest.TestCase):
214355
def test_primary_and_serial_validation_have_inner_and_outer_timeouts(self):
215356
with tempfile.TemporaryDirectory() as temporary_directory:

0 commit comments

Comments
 (0)