test: expand cli coverage and fix ruff issues

This commit is contained in:
vladkens
2026-05-21 21:31:26 +03:00
parent 59a658da1b
commit f8b9c95f97
13 changed files with 155 additions and 46 deletions
+2 -2
View File
@@ -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
View File
@@ -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"
+106
View File
@@ -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)
+1 -1
View File
@@ -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
+1 -1
View File
@@ -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):
+2 -2
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
View File
@@ -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
+1 -1
View File
@@ -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
View File
@@ -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}