1
0

generate_roles.py 3.7 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283848586878889
  1. import sys
  2. import os
  3. import json
  4. import argparse
  5. from typing import List
  6. # Ensure project root is in python path
  7. sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
  8. from sqlalchemy.orm import Session
  9. from app.db.session import SessionLocal
  10. from app.agent.real_god import RealGodAgent
  11. from app.crud import create_persona, get_user_by_username
  12. from app.schemas import PersonaCreate
  13. def generate_roles(topic: str, n: int = 3, owner_username: str = "admin"):
  14. """
  15. Generate roles using RealGodAgent and save to DB.
  16. """
  17. db = SessionLocal()
  18. try:
  19. # Get owner (admin)
  20. user = get_user_by_username(db, owner_username)
  21. if not user:
  22. print(f"User {owner_username} not found. Please create it first.")
  23. return []
  24. print(f"Generating {n} roles for topic: '{topic}'...")
  25. agent = RealGodAgent()
  26. generated_names = []
  27. created_persona_ids = []
  28. for i in range(n):
  29. print(f"Generating role {i+1}/{n}...")
  30. # Run agent for 1 persona
  31. # We collect the result from the generator
  32. for event in agent.run(topic, n=1, generated_names=generated_names):
  33. if event["type"] == "result":
  34. personas_data = event["content"]
  35. for p_data in personas_data:
  36. name = p_data.get('name')
  37. if name:
  38. generated_names.append(name)
  39. try:
  40. # Handle theories field
  41. if isinstance(p_data.get('theories'), str):
  42. try:
  43. p_data['theories'] = json.loads(p_data['theories'])
  44. except:
  45. p_data['theories'] = []
  46. # Create Schema
  47. persona_create = PersonaCreate(**p_data)
  48. persona_create.is_public = True # Make them public for experiments
  49. # Save to DB
  50. db_persona = create_persona(db=db, persona=persona_create, owner_id=user.id)
  51. created_persona_ids.append(db_persona.id)
  52. print(f" -> Created persona: {db_persona.name} (ID: {db_persona.id})")
  53. except Exception as e:
  54. print(f" -> Error saving persona: {e}")
  55. elif event["type"] == "thought":
  56. print(f" [Thought]: {event['content']}")
  57. elif event["type"] == "action":
  58. print(f" [Action]: {event['content']}")
  59. elif event["type"] == "observation":
  60. print(f" [Observation]: {event['content'][:100]}...")
  61. elif event["type"] == "error":
  62. print(f" [Error]: {event['content']}")
  63. print(f"\nGeneration Complete. Created Persona IDs: {created_persona_ids}")
  64. return created_persona_ids
  65. finally:
  66. db.close()
  67. if __name__ == "__main__":
  68. parser = argparse.ArgumentParser(description="Generate roles for a forum topic.")
  69. parser.add_argument("topic", type=str, help="The topic/theme for the roles.")
  70. parser.add_argument("--n", type=int, default=3, help="Number of roles to generate.")
  71. parser.add_argument("--owner", type=str, default="admin", help="Username of the owner.")
  72. args = parser.parse_args()
  73. generate_roles(args.topic, args.n, args.owner)