Something went wrong. Try again.
This repository has no description
Something went wrong. Try again.
123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273import argparseimport asynciofrom datetime import datetimeimport gzipimport jsonimport loggingimport osimport sysfrom typing import Tuple, List, Dict
from atproto import AsyncClientfrom atproto import exceptions as at_exceptionsfrom atproto_client.models.app.bsky.feed.defs import FeedViewPost
from utils import get_accounts, load_checkpoint, RateLimit, BSKY_API_LIMIT
logger = logging.getLogger(__name__)logger.setLevel(logging.INFO)
# Create formatterformatter = logging.Formatter( "%(asctime)s | %(levelname)s | %(message)s", datefmt="%Y-%m-%d %H:%M:%S")
# Console handlerconsole_handler = logging.StreamHandler(sys.stdout)console_handler.setFormatter(formatter)logger.addHandler(console_handler)
BATCH_SIZE = 10REQUIRED_ENV = ("BSKY_USER", "BSKY_APP_PW")
def process_post(top: FeedViewPost): post = top.post data = { "author": post.author.did, "text": post.record.text, "cid": post.cid, "created_at": post.record.created_at, "repost": False, } if ( top.reason is not None and top.reason.py_type == "app.bsky.feed.defs#reasonRepost" ): data["repost"] = True
if post.embed: data["embed"] = {} if post.embed.py_type == "app.bsky.embed.external#view": data["embed"] = { "title": post.embed.external.title, "description": post.embed.external.description, "uri": post.embed.external.uri, "thumb": post.embed.external.thumb, } elif post.embed.py_type == "app.bsky.embed.record#view": # Ignore everything thats not a quote-tweet if post.embed.record.py_type == "app.bsky.embed.record#viewRecord": data["embed"] = { "author": post.embed.record.author.did, "text": post.embed.record.value.text, "cid": post.embed.record.cid, "created_at": post.embed.record.value.created_at, } elif post.embed.py_type == "app.bsky.embed.images#view": data["embed"]["images"] = [] for image in post.embed.images: data["embed"]["images"].append( { "alt_text": image.alt, "full_url": image.fullsize, "thumb_url": image.thumb, } ) elif post.embed.py_type == "app.bsky.embed.video#view": data["embed"]["video"] = { "alt": post.embed.alt, "full_url": post.embed.playlist, "thumb_url": post.embed.thumbnail, }
if top.reply: if top.reply.parent.py_type == "app.bsky.feed.defs#postView": data["reply_parent"] = {} data["reply_parent"]["author"] = top.reply.parent.author.did data["reply_parent"]["text"] = top.reply.parent.record.text data["reply_parent"]["cid"] = top.reply.parent.cid data["reply_parent"]["created_at"] = top.reply.parent.record.created_at if top.reply.root.py_type == "app.bsky.feed.defs#postView": data["reply_parent"]["root_cid"] = top.reply.root.cid
return data
async def get_all_posts( client: AsyncClient, rate_limit: RateLimit, account_did: str, start_dt: datetime, end_dt: datetime,) -> Tuple[List[Dict], str]: posts: List[Dict] = [] await rate_limit.acquire() try: data = await client.get_author_feed( actor=account_did, filter="posts_and_author_threads", ) # If user can't be accessed just return an empty list to skip next time except at_exceptions.BadRequestError as e: if e.response.status_code == 400: return [], account_did else: logger.info(f"Error status code: {e.response.status_code}") raise e
for top in data.feed: dt = datetime.strptime(top.post.indexed_at, "%Y-%m-%dT%H:%M:%S.%fZ") if start_dt <= dt and dt < end_dt: parsed = process_post(top) if parsed is not None: posts.append(parsed)
hit_start_window = False while data.cursor and not hit_start_window: await rate_limit.acquire() data = await client.get_author_feed( actor=account_did, filter="posts_and_author_threads", cursor=data.cursor )
for top in data.feed: dt = datetime.strptime(top.post.indexed_at, "%Y-%m-%dT%H:%M:%S.%fZ") if start_dt <= dt and dt < end_dt: parsed = process_post(top) if parsed is not None: posts.append(parsed) if dt < start_dt: hit_start_window = True
return posts, account_did
async def retrieve_posts( user: str, app_pw: str, graph_file: str, checkpoint_dir: str, start_dt: datetime, end_dt: datetime,): # Checkpoint folders contain one file per user completed_accounts = load_checkpoint(checkpoint_dir) accts = get_accounts(graph_file, completed_accounts) logger.info(f"Num of accounts to retrieve posts from: {len(accts)}")
client = AsyncClient() await client.login(user, app_pw)
# Get all posts for accounts batch_count = 0 fail_count = 0 rate_limiter = RateLimit(BSKY_API_LIMIT) for i in range(0, len(accts), BATCH_SIZE): batch = [acct for acct, _ in accts[i : i + BATCH_SIZE]] for result in asyncio.as_completed( [ get_all_posts(client, rate_limiter, did, start_dt, end_dt) for did in batch ] ): try: posts, did = await result # Save posts with gzip.open( os.path.join(checkpoint_dir, did + ".gz"), "wt" ) as out_file: for post in posts: out_file.write(json.dumps(post) + "\n") except at_exceptions.BadRequestError as e: # Bad request is probably a profile that's private or deleted logger.info(f"Bad Request: {e.response.content.error}") continue except Exception as e: logger.error(f"Failed to get posts: {e}", exc_info=1) fail_count += 1 if fail_count >= 100: logger.error("Hitting error threshold, exiting...") sys.exit(1) continue
batch_count += 1 if batch_count % 10 == 0: logger.info(f"Completed batch: {batch_count}")
def main(): for key in REQUIRED_ENV: if key not in os.environ: raise ValueError(f"Must set '{key}' env var")
user_name = os.environ["BSKY_USER"] app_pw = os.environ["BSKY_APP_PW"]
parser = argparse.ArgumentParser( prog="GetPosts", description="Get all posts for accounts in provided follow graph", ) parser.add_argument( "--graph-file", dest="graph_file", required=True, help="File with follow graph", ) parser.add_argument( "--save-dir", dest="save_dir", required=True, help="Where to store crawl data", ) parser.add_argument( "--start", dest="start", required=True, help="Date to start saving posts from (YYYY-MM-DD)", ) parser.add_argument( "--end", dest="end", required=True, help="Date to stop (exclusive) saving posts from (YYYY-MM-DD)", ) args = parser.parse_args()
if args.save_dir is None and args.ckpt is None: logger.error("Must provide save dir or checkpoint dir") sys.exit(1)
try: start = datetime.strptime(args.start, "%Y-%m-%d") except: logger.error("Invalid start date") sys.exit(1)
try: end = datetime.strptime(args.end, "%Y-%m-%d") except: logger.error("Invalid end date") sys.exit(1)
if end <= start: logger.error( "Start date has to be before date, what're you trying to do man..." ) sys.exit(1)
asyncio.run( retrieve_posts( user_name, app_pw, graph_file=args.graph_file, checkpoint_dir=args.save_dir, start_dt=start, end_dt=end, ) )
if __name__ == "__main__": main()