mirror of
https://github.com/vladkens/twscrape.git
synced 2026-10-10 15:17:19 -04:00
test: expand cli coverage and fix ruff issues
This commit is contained in:
@@ -11,9 +11,9 @@ async def download_file(client: httpx.AsyncClient, url: str, outdir: str):
|
||||
outpath = os.path.join(outdir, filename)
|
||||
|
||||
async with client.stream("GET", url) as resp:
|
||||
with open(outpath, "wb") as f:
|
||||
with open(outpath, "wb") as fp:
|
||||
async for chunk in resp.aiter_bytes():
|
||||
f.write(chunk)
|
||||
fp.write(chunk)
|
||||
|
||||
|
||||
async def load_user_media(api: API, user_id: int, outdir: str):
|
||||
|
||||
+10
-5
@@ -1,7 +1,3 @@
|
||||
[build-system]
|
||||
requires = ["hatchling>=1.9.1"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[project]
|
||||
name = "twscrape"
|
||||
version = "0.17.0"
|
||||
@@ -45,6 +41,10 @@ repository = "https://github.com/vladkens/twscrape"
|
||||
[project.scripts]
|
||||
twscrape = "twscrape.cli:run"
|
||||
|
||||
[build-system]
|
||||
requires = ["hatchling>=1.9.1"]
|
||||
build-backend = "hatchling.build"
|
||||
|
||||
[tool.hatch.build.targets.wheel]
|
||||
packages = ["twscrape"]
|
||||
|
||||
@@ -62,4 +62,9 @@ line-length = 99
|
||||
target-version = "py310"
|
||||
|
||||
[tool.ruff.lint]
|
||||
ignore = ["E501"]
|
||||
select = ["E", "F", "I", "UP", "C4", "SIM"]
|
||||
ignore = ["E501", "UP035", "SIM105"]
|
||||
|
||||
[tool.pyright]
|
||||
include = ["apps", "clients", "strategy", "lib", "scripts"]
|
||||
typeCheckingMode = "standard"
|
||||
|
||||
@@ -1,5 +1,7 @@
|
||||
import argparse
|
||||
import json
|
||||
|
||||
from tests.test_parser import fake_rep
|
||||
from twscrape import cli
|
||||
|
||||
|
||||
@@ -58,3 +60,107 @@ async def test_add_cookie_prompts_securely_when_missing(tmp_path, monkeypatch):
|
||||
"username": "user1",
|
||||
"cookies": "auth_token=prompted; ct0=csrf",
|
||||
}
|
||||
|
||||
|
||||
async def test_add_accounts_prints_next_step(tmp_path, monkeypatch, capsys):
|
||||
called = {}
|
||||
|
||||
async def mock_load_from_file(self, file_path, line_format):
|
||||
called["file_path"] = file_path
|
||||
called["line_format"] = line_format
|
||||
|
||||
monkeypatch.setattr(cli.AccountsPool, "load_from_file", mock_load_from_file)
|
||||
|
||||
args = argparse.Namespace(
|
||||
command="add_accounts",
|
||||
debug=False,
|
||||
db=str(tmp_path / "test.db"),
|
||||
email_first=False,
|
||||
manual=False,
|
||||
file_path="accounts.txt",
|
||||
line_format="username:password:email:email_password",
|
||||
)
|
||||
|
||||
await cli.main(args)
|
||||
|
||||
out = capsys.readouterr().out
|
||||
assert called == {
|
||||
"file_path": "accounts.txt",
|
||||
"line_format": "username:password:email:email_password",
|
||||
}
|
||||
assert "twscrape login_accounts" in out
|
||||
|
||||
|
||||
async def test_search_prints_parsed_tweets(tmp_path, monkeypatch, capsys):
|
||||
async def mock_search_raw(self, q, limit=-1, kv=None):
|
||||
yield fake_rep("raw_search")
|
||||
|
||||
monkeypatch.setattr(cli.API, "search_raw", mock_search_raw)
|
||||
|
||||
args = argparse.Namespace(
|
||||
command="search",
|
||||
debug=False,
|
||||
db=str(tmp_path / "test.db"),
|
||||
email_first=False,
|
||||
manual=False,
|
||||
raw=False,
|
||||
arg_name="query",
|
||||
query="elon musk lang:en",
|
||||
limit=1,
|
||||
)
|
||||
|
||||
await cli.main(args)
|
||||
|
||||
out = capsys.readouterr().out.strip().splitlines()
|
||||
assert len(out) > 0
|
||||
doc = json.loads(out[0])
|
||||
assert isinstance(doc["id"], int)
|
||||
assert doc["user"]["username"] is not None
|
||||
|
||||
|
||||
async def test_user_by_login_prints_parsed_user(tmp_path, monkeypatch, capsys):
|
||||
async def mock_user_by_login_raw(self, login, kv=None):
|
||||
return fake_rep("raw_user_by_login")
|
||||
|
||||
monkeypatch.setattr(cli.API, "user_by_login_raw", mock_user_by_login_raw)
|
||||
|
||||
args = argparse.Namespace(
|
||||
command="user_by_login",
|
||||
debug=False,
|
||||
db=str(tmp_path / "test.db"),
|
||||
email_first=False,
|
||||
manual=False,
|
||||
raw=False,
|
||||
arg_name="username",
|
||||
username="xdevelopers",
|
||||
)
|
||||
|
||||
await cli.main(args)
|
||||
|
||||
doc = json.loads(capsys.readouterr().out.strip())
|
||||
assert doc["id"] == 2244994945
|
||||
assert doc["username"] == "XDevelopers"
|
||||
|
||||
|
||||
async def test_tweet_details_raw_prints_raw_json(tmp_path, monkeypatch, capsys):
|
||||
async def mock_tweet_details_raw(self, twid, kv=None):
|
||||
return fake_rep("raw_tweet_details")
|
||||
|
||||
monkeypatch.setattr(cli.API, "tweet_details_raw", mock_tweet_details_raw)
|
||||
|
||||
args = argparse.Namespace(
|
||||
command="tweet_details",
|
||||
debug=False,
|
||||
db=str(tmp_path / "test.db"),
|
||||
email_first=False,
|
||||
manual=False,
|
||||
raw=True,
|
||||
arg_name="tweet_id",
|
||||
tweet_id=1649191520250245121,
|
||||
)
|
||||
|
||||
await cli.main(args)
|
||||
|
||||
doc = json.loads(capsys.readouterr().out.strip())
|
||||
assert "data" in doc
|
||||
assert "threaded_conversation_with_injections_v2" in json.dumps(doc)
|
||||
|
||||
@@ -512,7 +512,7 @@ async def test_issue_56():
|
||||
raw = fake_rep("_issue_56").json()
|
||||
doc = parse_tweet(raw, 1682072224013099008)
|
||||
assert doc is not None
|
||||
assert len(set([x.tcourl for x in doc.links])) == len(doc.links)
|
||||
assert len({x.tcourl for x in doc.links}) == len(doc.links)
|
||||
assert len(doc.links) == 5
|
||||
|
||||
|
||||
|
||||
@@ -13,7 +13,7 @@ CF = tuple[AccountsPool, QueueClient]
|
||||
|
||||
async def get_locked(pool: AccountsPool) -> set[str]:
|
||||
rep = await pool.get_all()
|
||||
return set([x.username for x in rep if x.locks.get("SearchTimeline", None) is not None])
|
||||
return {x.username for x in rep if x.locks.get("SearchTimeline", None) is not None}
|
||||
|
||||
|
||||
async def test_lock_account_when_used(httpx_mock: HTTPXMock, client_fixture):
|
||||
|
||||
@@ -50,12 +50,12 @@ class AccountsPool:
|
||||
line_delim = guess_delim(line_format)
|
||||
tokens = line_format.split(line_delim)
|
||||
|
||||
required = set(["username", "password", "email", "email_password"])
|
||||
required = {"username", "password", "email", "email_password"}
|
||||
if not required.issubset(tokens):
|
||||
raise ValueError(f"Invalid line format: {line_format}")
|
||||
|
||||
accounts = []
|
||||
with open(filepath, "r") as f:
|
||||
with open(filepath) as f:
|
||||
lines = f.read().split("\n")
|
||||
lines = [x.strip() for x in lines if x.strip()]
|
||||
|
||||
|
||||
+6
-6
@@ -174,19 +174,19 @@ def run():
|
||||
|
||||
add_accounts = subparsers.add_parser("add_accounts", help="Add accounts from file")
|
||||
add_accounts.add_argument("file_path", help="File with accounts")
|
||||
add_accounts.add_argument("line_format", help="args of Pool.add_account splited by same delim")
|
||||
add_accounts.add_argument("line_format", help="Account fields separated by delimiter")
|
||||
|
||||
add_cookie = subparsers.add_parser("add_cookie", help="Add account with cookies")
|
||||
add_cookie = subparsers.add_parser("add_cookie", help="Add one account from cookies")
|
||||
add_cookie.add_argument("username", help="Twitter/X username")
|
||||
add_cookie.add_argument("cookies", nargs="?", default=None, help="Cookie string")
|
||||
|
||||
del_accounts = subparsers.add_parser("del_accounts", help="Delete accounts")
|
||||
del_accounts = subparsers.add_parser("del_accounts", help="Delete accounts by username")
|
||||
del_accounts.add_argument("usernames", nargs="+", default=[], help="Usernames to delete")
|
||||
|
||||
login_cmd = subparsers.add_parser("login_accounts", help="Login accounts")
|
||||
relogin = subparsers.add_parser("relogin", help="Re-login selected accounts")
|
||||
login_cmd = subparsers.add_parser("login_accounts", help="Log in inactive accounts")
|
||||
relogin = subparsers.add_parser("relogin", help="Re-log in selected accounts")
|
||||
relogin.add_argument("usernames", nargs="+", default=[], help="Usernames to re-login")
|
||||
re_failed = subparsers.add_parser("relogin_failed", help="Retry login for failed accounts")
|
||||
re_failed = subparsers.add_parser("relogin_failed", help="Retry failed account logins")
|
||||
|
||||
login_commands = [login_cmd, relogin, re_failed]
|
||||
for cmd in login_commands:
|
||||
|
||||
+7
-9
@@ -34,7 +34,7 @@ def lock_retry(max_retries=10):
|
||||
|
||||
|
||||
async def get_sqlite_version():
|
||||
async with aiosqlite.connect(":memory:") as db:
|
||||
async with aiosqlite.connect(":memory:") as db: # noqa: SIM117
|
||||
async with db.execute("SELECT SQLITE_VERSION()") as cur:
|
||||
rs = await cur.fetchone()
|
||||
return rs[0] if rs else "3.0.0"
|
||||
@@ -138,18 +138,16 @@ async def execute(db_path: str, qs: str, params: dict | None = None):
|
||||
|
||||
@lock_retry()
|
||||
async def fetchone(db_path: str, qs: str, params: dict | None = None):
|
||||
async with DB(db_path) as db:
|
||||
async with db.execute(qs, params) as cur:
|
||||
row = await cur.fetchone()
|
||||
return row
|
||||
async with DB(db_path) as db, db.execute(qs, params) as cur:
|
||||
row = await cur.fetchone()
|
||||
return row
|
||||
|
||||
|
||||
@lock_retry()
|
||||
async def fetchall(db_path: str, qs: str, params: dict | None = None):
|
||||
async with DB(db_path) as db:
|
||||
async with db.execute(qs, params) as cur:
|
||||
rows = await cur.fetchall()
|
||||
return rows
|
||||
async with DB(db_path) as db, db.execute(qs, params) as cur:
|
||||
rows = await cur.fetchall()
|
||||
return rows
|
||||
|
||||
|
||||
@lock_retry()
|
||||
|
||||
+1
-1
@@ -87,7 +87,7 @@ async def imap_get_email_code(
|
||||
if code is not None:
|
||||
return code
|
||||
|
||||
if TWS_WAIT_EMAIL_CODE < time.time() - start_time:
|
||||
if time.time() - start_time > TWS_WAIT_EMAIL_CODE:
|
||||
raise EmailCodeTimeoutError(f"Email code timeout ({TWS_WAIT_EMAIL_CODE} sec)")
|
||||
|
||||
await asyncio.sleep(5)
|
||||
|
||||
+2
-2
@@ -274,6 +274,6 @@ async def login(acc: Account, cfg: LoginConfig | None = None) -> Account:
|
||||
client.headers["x-twitter-auth-type"] = "OAuth2Session"
|
||||
|
||||
acc.active = True
|
||||
acc.headers = {k: v for k, v in client.headers.items()}
|
||||
acc.cookies = {k: v for k, v in client.cookies.items()}
|
||||
acc.headers = dict(client.headers.items())
|
||||
acc.cookies = dict(client.cookies.items())
|
||||
return acc
|
||||
|
||||
+14
-14
@@ -70,9 +70,9 @@ class TextLink(JSONTrait):
|
||||
|
||||
@staticmethod
|
||||
def parse(obj: dict):
|
||||
url1 = obj.get("expanded_url", None)
|
||||
url2 = obj.get("url", None)
|
||||
text = obj.get("display_url", None)
|
||||
url1 = obj.get("expanded_url")
|
||||
url2 = obj.get("url")
|
||||
text = obj.get("display_url")
|
||||
|
||||
if not isinstance(url1, str) or not isinstance(url2, str):
|
||||
return None
|
||||
@@ -266,8 +266,8 @@ class Tweet(JSONTrait):
|
||||
viewCount: int | None = None
|
||||
retweetedTweet: Optional["Tweet"] = None
|
||||
quotedTweet: Optional["Tweet"] = None
|
||||
place: Optional[Place] = None
|
||||
coordinates: Optional[Coordinates] = None
|
||||
place: Place | None = None
|
||||
coordinates: Coordinates | None = None
|
||||
inReplyToTweetId: int | None = None
|
||||
inReplyToTweetIdStr: str | None = None
|
||||
inReplyToUser: UserRef | None = None
|
||||
@@ -331,12 +331,12 @@ class Tweet(JSONTrait):
|
||||
inReplyToTweetId=int_or(obj, "in_reply_to_status_id_str"),
|
||||
inReplyToTweetIdStr=get_or(obj, "in_reply_to_status_id_str"),
|
||||
inReplyToUser=_get_reply_user(obj, res),
|
||||
source=obj.get("source", None),
|
||||
source=obj.get("source"),
|
||||
sourceUrl=_get_source_url(obj),
|
||||
sourceLabel=_get_source_label(obj),
|
||||
media=Media.parse(obj),
|
||||
card=_parse_card(obj, url),
|
||||
possibly_sensitive=obj.get("possibly_sensitive", None),
|
||||
possibly_sensitive=obj.get("possibly_sensitive"),
|
||||
)
|
||||
|
||||
# issue #42 – restore full rt text
|
||||
@@ -513,8 +513,8 @@ class TrendUrl(JSONTrait):
|
||||
|
||||
@dataclass
|
||||
class TrendMetadata(JSONTrait):
|
||||
domain_context: Optional[str]
|
||||
meta_description: Optional[str]
|
||||
domain_context: str | None
|
||||
meta_description: str | None
|
||||
url: TrendUrl
|
||||
|
||||
@staticmethod
|
||||
@@ -538,8 +538,8 @@ class GroupedTrend(JSONTrait):
|
||||
|
||||
@dataclass
|
||||
class Trend(JSONTrait):
|
||||
id: Optional[str]
|
||||
rank: Optional[str | int]
|
||||
id: str | None
|
||||
rank: str | int | None
|
||||
name: str
|
||||
trend_url: TrendUrl
|
||||
trend_metadata: TrendMetadata
|
||||
@@ -710,7 +710,7 @@ def _parse_card(obj: dict, url: str):
|
||||
|
||||
|
||||
def _get_reply_user(tw_obj: dict, res: dict):
|
||||
user_id = tw_obj.get("in_reply_to_user_id_str", None)
|
||||
user_id = tw_obj.get("in_reply_to_user_id_str")
|
||||
if user_id is None:
|
||||
return None
|
||||
|
||||
@@ -727,14 +727,14 @@ def _get_reply_user(tw_obj: dict, res: dict):
|
||||
|
||||
|
||||
def _get_source_url(tw_obj: dict):
|
||||
source = tw_obj.get("source", None)
|
||||
source = tw_obj.get("source")
|
||||
if source and (match := re.search(r'href=[\'"]?([^\'" >]+)', source)):
|
||||
return str(match.group(1))
|
||||
return None
|
||||
|
||||
|
||||
def _get_source_label(tw_obj: dict):
|
||||
source = tw_obj.get("source", None)
|
||||
source = tw_obj.get("source")
|
||||
if source and (match := re.search(r">([^<]*)<", source)):
|
||||
return str(match.group(1))
|
||||
return None
|
||||
|
||||
@@ -186,7 +186,7 @@ class QueueClient:
|
||||
|
||||
err_msg = "OK"
|
||||
if "errors" in res:
|
||||
err_msg = set([f"({x.get('code', -1)}) {x['message']}" for x in res["errors"]])
|
||||
err_msg = {f"({x.get('code', -1)}) {x['message']}" for x in res["errors"]}
|
||||
err_msg = "; ".join(list(err_msg))
|
||||
|
||||
log_msg = f"{rep.status_code:3d} - {req_id(rep)} - {err_msg}"
|
||||
|
||||
+2
-2
@@ -109,7 +109,7 @@ def find_obj(obj: dict, fn: Callable[[dict], bool]) -> Any | None:
|
||||
|
||||
|
||||
def get_typed_object(obj: dict, res: defaultdict[str, list]):
|
||||
obj_type = obj.get("__typename", None)
|
||||
obj_type = obj.get("__typename")
|
||||
if obj_type is not None:
|
||||
res[obj_type].append(obj)
|
||||
|
||||
@@ -271,7 +271,7 @@ def to_old_rep(obj: dict) -> dict[str, dict]:
|
||||
if res := _to_old_user(x):
|
||||
users[str(res["id_str"])] = res
|
||||
|
||||
trends = [x for x in tmp.get("TimelineTrend", [])]
|
||||
trends = list(tmp.get("TimelineTrend", []))
|
||||
trends = {x["name"]: x for x in trends}
|
||||
|
||||
return {"tweets": {**tw1, **tw2}, "users": users, "trends": trends}
|
||||
|
||||
Reference in New Issue
Block a user