From 1096fbce9dec1592173a98e670648b54fb7a9722 Mon Sep 17 00:00:00 2001 From: W11 Date: Sat, 19 Sep 2026 18:22:52 +0800 Subject: [PATCH] =?UTF-8?q?init:=20=E8=87=AA=20zomaintain/backend/rdplib?= =?UTF-8?q?=20=E5=B9=B3=E7=A7=BB=E7=8B=AC=E7=AB=8B=E6=88=90=E5=BA=93;=20mo?= =?UTF-8?q?dule=20path=20=E4=BB=8E=E4=B8=8A=E6=B8=B8=20github.com/nakagami?= =?UTF-8?q?/grdp=20=E6=94=B9=E4=B8=BA=20git.zeroonesoft.cn/golib/rdplib?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- LICENSE | 674 ++++++ README.md | 119 + bmpcache_test.go | 105 + convert_amd64.go | 55 + convert_amd64.s | 210 ++ convert_arm64.go | 52 + convert_arm64.s | 171 ++ convert_generic.go | 33 + core/io.go | 131 + core/io_test.go | 17 + core/log.go | 20 + core/mppc.go | 205 ++ core/mppc_test.go | 353 +++ core/rle.go | 997 ++++++++ core/rle_test.go | 26 + core/socket.go | 127 + core/types.go | 26 + core/util.go | 53 + doc.go | 36 + emission/emitter.go | 304 +++ emission/emitter_test.go | 46 + go.mod | 8 + go.sum | 4 + grdp.go | 2069 ++++++++++++++++ grdp_scancode_test.go | 33 + plugin/addins.go | 308 +++ plugin/channel.go | 308 +++ plugin/cliprdr/cliprdr.go | 265 ++ plugin/cliprdr/cliprdr_generic.go | 181 ++ plugin/cliprdr/cliprdr_image_test.go | 198 ++ plugin/cliprdr/cliprdr_test.go | 44 + plugin/cliprdr/cliprdr_types.go | 327 +++ plugin/cliprdr/file_clip.go | 479 ++++ plugin/cliprdr/file_clip_test.go | 527 ++++ plugin/cliprdr/handler.go | 678 ++++++ plugin/cliprdr/html_format.go | 77 + plugin/cliprdr/html_format_test.go | 104 + plugin/cliprdr/image_dib.go | 114 + plugin/drdynvc/dvc.go | 364 +++ plugin/rail/rail.go | 452 ++++ plugin/rdpdr/drive.go | 256 ++ plugin/rdpdr/rdpdr.go | 695 ++++++ plugin/rdpdr/rdpdr_test.go | 428 ++++ plugin/rdpedisp/rdpedisp.go | 186 ++ plugin/rdpedisp/rdpedisp_test.go | 63 + plugin/rdpgfx/avc.go | 2329 ++++++++++++++++++ plugin/rdpgfx/avc_test.go | 33 + plugin/rdpgfx/clear.go | 710 ++++++ plugin/rdpgfx/clear_test.go | 295 +++ plugin/rdpgfx/convert_parallel_test.go | 229 ++ plugin/rdpgfx/ffmpeg/h264_ffmpeg.go | 3077 ++++++++++++++++++++++++ plugin/rdpgfx/h264_decoder.go | 156 ++ plugin/rdpgfx/h264_scan.go | 67 + plugin/rdpgfx/ict_arm64.go | 37 + plugin/rdpgfx/ict_arm64.s | 153 ++ plugin/rdpgfx/ict_generic.go | 38 + plugin/rdpgfx/rdpgfx.go | 2816 ++++++++++++++++++++++ plugin/rdpgfx/rdpgfx_cache_test.go | 234 ++ plugin/rdpgfx/rfx.go | 321 +++ plugin/rdpgfx/rfx_dwt_shared.go | 221 ++ plugin/rdpgfx/rfx_pool.go | 50 + plugin/rdpgfx/rfx_progressive.go | 1321 ++++++++++ plugin/rdpgfx/rfx_progressive_test.go | 57 + plugin/rdpgfx/rfx_rlgr.go | 452 ++++ plugin/rdpgfx/rfx_rlgr_test.go | 581 +++++ plugin/rdpgfx/zgfx.go | 493 ++++ plugin/rdpsnd/aac/decoder.go | 18 + plugin/rdpsnd/aac/decoder_darwin.go | 210 ++ plugin/rdpsnd/aac/decoder_stub.go | 22 + plugin/rdpsnd/rdpsnd.go | 524 ++++ protocol/lic/lic.go | 183 ++ protocol/nla/cssp.go | 98 + protocol/nla/cssp_test.go | 16 + protocol/nla/encode.go | 46 + protocol/nla/encode_test.go | 32 + protocol/nla/ntlm.go | 530 ++++ protocol/nla/ntlm_test.go | 53 + protocol/pdu/caps.go | 857 +++++++ protocol/pdu/data.go | 2337 ++++++++++++++++++ protocol/pdu/data_test.go | 327 +++ protocol/pdu/orders.go | 1243 ++++++++++ protocol/pdu/pdu.go | 935 +++++++ protocol/pdu/ycbcr_amd64.go | 34 + protocol/pdu/ycbcr_amd64.s | 100 + protocol/pdu/ycbcr_arm64.go | 31 + protocol/pdu/ycbcr_arm64.s | 79 + protocol/pdu/ycbcr_generic.go | 55 + protocol/sec/sec.go | 994 ++++++++ protocol/t125/ber/ber.go | 189 ++ protocol/t125/gcc/gcc.go | 671 ++++++ protocol/t125/mcs.go | 753 ++++++ protocol/t125/per/per.go | 211 ++ protocol/tpkt/tpkt.go | 293 +++ protocol/x224/x224.go | 351 +++ 94 files changed, 36790 insertions(+) create mode 100644 LICENSE create mode 100644 README.md create mode 100644 bmpcache_test.go create mode 100644 convert_amd64.go create mode 100644 convert_amd64.s create mode 100644 convert_arm64.go create mode 100644 convert_arm64.s create mode 100644 convert_generic.go create mode 100644 core/io.go create mode 100644 core/io_test.go create mode 100644 core/log.go create mode 100644 core/mppc.go create mode 100644 core/mppc_test.go create mode 100644 core/rle.go create mode 100644 core/rle_test.go create mode 100644 core/socket.go create mode 100644 core/types.go create mode 100644 core/util.go create mode 100644 doc.go create mode 100644 emission/emitter.go create mode 100644 emission/emitter_test.go create mode 100644 go.mod create mode 100644 go.sum create mode 100644 grdp.go create mode 100644 grdp_scancode_test.go create mode 100644 plugin/addins.go create mode 100644 plugin/channel.go create mode 100644 plugin/cliprdr/cliprdr.go create mode 100644 plugin/cliprdr/cliprdr_generic.go create mode 100644 plugin/cliprdr/cliprdr_image_test.go create mode 100644 plugin/cliprdr/cliprdr_test.go create mode 100644 plugin/cliprdr/cliprdr_types.go create mode 100644 plugin/cliprdr/file_clip.go create mode 100644 plugin/cliprdr/file_clip_test.go create mode 100644 plugin/cliprdr/handler.go create mode 100644 plugin/cliprdr/html_format.go create mode 100644 plugin/cliprdr/html_format_test.go create mode 100644 plugin/cliprdr/image_dib.go create mode 100644 plugin/drdynvc/dvc.go create mode 100644 plugin/rail/rail.go create mode 100644 plugin/rdpdr/drive.go create mode 100644 plugin/rdpdr/rdpdr.go create mode 100644 plugin/rdpdr/rdpdr_test.go create mode 100644 plugin/rdpedisp/rdpedisp.go create mode 100644 plugin/rdpedisp/rdpedisp_test.go create mode 100644 plugin/rdpgfx/avc.go create mode 100644 plugin/rdpgfx/avc_test.go create mode 100644 plugin/rdpgfx/clear.go create mode 100644 plugin/rdpgfx/clear_test.go create mode 100644 plugin/rdpgfx/convert_parallel_test.go create mode 100644 plugin/rdpgfx/ffmpeg/h264_ffmpeg.go create mode 100644 plugin/rdpgfx/h264_decoder.go create mode 100644 plugin/rdpgfx/h264_scan.go create mode 100644 plugin/rdpgfx/ict_arm64.go create mode 100644 plugin/rdpgfx/ict_arm64.s create mode 100644 plugin/rdpgfx/ict_generic.go create mode 100644 plugin/rdpgfx/rdpgfx.go create mode 100644 plugin/rdpgfx/rdpgfx_cache_test.go create mode 100644 plugin/rdpgfx/rfx.go create mode 100644 plugin/rdpgfx/rfx_dwt_shared.go create mode 100644 plugin/rdpgfx/rfx_pool.go create mode 100644 plugin/rdpgfx/rfx_progressive.go create mode 100644 plugin/rdpgfx/rfx_progressive_test.go create mode 100644 plugin/rdpgfx/rfx_rlgr.go create mode 100644 plugin/rdpgfx/rfx_rlgr_test.go create mode 100644 plugin/rdpgfx/zgfx.go create mode 100644 plugin/rdpsnd/aac/decoder.go create mode 100644 plugin/rdpsnd/aac/decoder_darwin.go create mode 100644 plugin/rdpsnd/aac/decoder_stub.go create mode 100644 plugin/rdpsnd/rdpsnd.go create mode 100644 protocol/lic/lic.go create mode 100644 protocol/nla/cssp.go create mode 100644 protocol/nla/cssp_test.go create mode 100644 protocol/nla/encode.go create mode 100644 protocol/nla/encode_test.go create mode 100644 protocol/nla/ntlm.go create mode 100644 protocol/nla/ntlm_test.go create mode 100644 protocol/pdu/caps.go create mode 100644 protocol/pdu/data.go create mode 100644 protocol/pdu/data_test.go create mode 100644 protocol/pdu/orders.go create mode 100644 protocol/pdu/pdu.go create mode 100644 protocol/pdu/ycbcr_amd64.go create mode 100644 protocol/pdu/ycbcr_amd64.s create mode 100644 protocol/pdu/ycbcr_arm64.go create mode 100644 protocol/pdu/ycbcr_arm64.s create mode 100644 protocol/pdu/ycbcr_generic.go create mode 100644 protocol/sec/sec.go create mode 100644 protocol/t125/ber/ber.go create mode 100644 protocol/t125/gcc/gcc.go create mode 100644 protocol/t125/mcs.go create mode 100644 protocol/t125/per/per.go create mode 100644 protocol/tpkt/tpkt.go create mode 100644 protocol/x224/x224.go diff --git a/LICENSE b/LICENSE new file mode 100644 index 0000000..20d40b6 --- /dev/null +++ b/LICENSE @@ -0,0 +1,674 @@ + GNU GENERAL PUBLIC LICENSE + Version 3, 29 June 2007 + + Copyright (C) 2007 Free Software Foundation, Inc. + Everyone is permitted to copy and distribute verbatim copies + of this license document, but changing it is not allowed. + + Preamble + + The GNU General Public License is a free, copyleft license for +software and other kinds of works. + + The licenses for most software and other practical works are designed +to take away your freedom to share and change the works. By contrast, +the GNU General Public License is intended to guarantee your freedom to +share and change all versions of a program--to make sure it remains free +software for all its users. We, the Free Software Foundation, use the +GNU General Public License for most of our software; it applies also to +any other work released this way by its authors. You can apply it to +your programs, too. + + When we speak of free software, we are referring to freedom, not +price. Our General Public Licenses are designed to make sure that you +have the freedom to distribute copies of free software (and charge for +them if you wish), that you receive source code or can get it if you +want it, that you can change the software or use pieces of it in new +free programs, and that you know you can do these things. + + To protect your rights, we need to prevent others from denying you +these rights or asking you to surrender the rights. Therefore, you have +certain responsibilities if you distribute copies of the software, or if +you modify it: responsibilities to respect the freedom of others. + + For example, if you distribute copies of such a program, whether +gratis or for a fee, you must pass on to the recipients the same +freedoms that you received. You must make sure that they, too, receive +or can get the source code. And you must show them these terms so they +know their rights. + + Developers that use the GNU GPL protect your rights with two steps: +(1) assert copyright on the software, and (2) offer you this License +giving you legal permission to copy, distribute and/or modify it. + + For the developers' and authors' protection, the GPL clearly explains +that there is no warranty for this free software. For both users' and +authors' sake, the GPL requires that modified versions be marked as +changed, so that their problems will not be attributed erroneously to +authors of previous versions. + + Some devices are designed to deny users access to install or run +modified versions of the software inside them, although the manufacturer +can do so. This is fundamentally incompatible with the aim of +protecting users' freedom to change the software. The systematic +pattern of such abuse occurs in the area of products for individuals to +use, which is precisely where it is most unacceptable. Therefore, we +have designed this version of the GPL to prohibit the practice for those +products. If such problems arise substantially in other domains, we +stand ready to extend this provision to those domains in future versions +of the GPL, as needed to protect the freedom of users. + + Finally, every program is threatened constantly by software patents. +States should not allow patents to restrict development and use of +software on general-purpose computers, but in those that do, we wish to +avoid the special danger that patents applied to a free program could +make it effectively proprietary. To prevent this, the GPL assures that +patents cannot be used to render the program non-free. + + The precise terms and conditions for copying, distribution and +modification follow. + + TERMS AND CONDITIONS + + 0. Definitions. + + "This License" refers to version 3 of the GNU General Public License. + + "Copyright" also means copyright-like laws that apply to other kinds of +works, such as semiconductor masks. + + "The Program" refers to any copyrightable work licensed under this +License. Each licensee is addressed as "you". "Licensees" and +"recipients" may be individuals or organizations. + + To "modify" a work means to copy from or adapt all or part of the work +in a fashion requiring copyright permission, other than the making of an +exact copy. The resulting work is called a "modified version" of the +earlier work or a work "based on" the earlier work. + + A "covered work" means either the unmodified Program or a work based +on the Program. + + To "propagate" a work means to do anything with it that, without +permission, would make you directly or secondarily liable for +infringement under applicable copyright law, except executing it on a +computer or modifying a private copy. Propagation includes copying, +distribution (with or without modification), making available to the +public, and in some countries other activities as well. + + To "convey" a work means any kind of propagation that enables other +parties to make or receive copies. Mere interaction with a user through +a computer network, with no transfer of a copy, is not conveying. + + An interactive user interface displays "Appropriate Legal Notices" +to the extent that it includes a convenient and prominently visible +feature that (1) displays an appropriate copyright notice, and (2) +tells the user that there is no warranty for the work (except to the +extent that warranties are provided), that licensees may convey the +work under this License, and how to view a copy of this License. If +the interface presents a list of user commands or options, such as a +menu, a prominent item in the list meets this criterion. + + 1. Source Code. + + The "source code" for a work means the preferred form of the work +for making modifications to it. "Object code" means any non-source +form of a work. + + A "Standard Interface" means an interface that either is an official +standard defined by a recognized standards body, or, in the case of +interfaces specified for a particular programming language, one that +is widely used among developers working in that language. + + The "System Libraries" of an executable work include anything, other +than the work as a whole, that (a) is included in the normal form of +packaging a Major Component, but which is not part of that Major +Component, and (b) serves only to enable use of the work with that +Major Component, or to implement a Standard Interface for which an +implementation is available to the public in source code form. A +"Major Component", in this context, means a major essential component +(kernel, window system, and so on) of the specific operating system +(if any) on which the executable work runs, or a compiler used to +produce the work, or an object code interpreter used to run it. + + The "Corresponding Source" for a work in object code form means all +the source code needed to generate, install, and (for an executable +work) run the object code and to modify the work, including scripts to +control those activities. However, it does not include the work's +System Libraries, or general-purpose tools or generally available free +programs which are used unmodified in performing those activities but +which are not part of the work. For example, Corresponding Source +includes interface definition files associated with source files for +the work, and the source code for shared libraries and dynamically +linked subprograms that the work is specifically designed to require, +such as by intimate data communication or control flow between those +subprograms and other parts of the work. + + The Corresponding Source need not include anything that users +can regenerate automatically from other parts of the Corresponding +Source. + + The Corresponding Source for a work in source code form is that +same work. + + 2. Basic Permissions. + + All rights granted under this License are granted for the term of +copyright on the Program, and are irrevocable provided the stated +conditions are met. This License explicitly affirms your unlimited +permission to run the unmodified Program. The output from running a +covered work is covered by this License only if the output, given its +content, constitutes a covered work. This License acknowledges your +rights of fair use or other equivalent, as provided by copyright law. + + You may make, run and propagate covered works that you do not +convey, without conditions so long as your license otherwise remains +in force. You may convey covered works to others for the sole purpose +of having them make modifications exclusively for you, or provide you +with facilities for running those works, provided that you comply with +the terms of this License in conveying all material for which you do +not control copyright. Those thus making or running the covered works +for you must do so exclusively on your behalf, under your direction +and control, on terms that prohibit them from making any copies of +your copyrighted material outside their relationship with you. + + Conveying under any other circumstances is permitted solely under +the conditions stated below. Sublicensing is not allowed; section 10 +makes it unnecessary. + + 3. Protecting Users' Legal Rights From Anti-Circumvention Law. + + No covered work shall be deemed part of an effective technological +measure under any applicable law fulfilling obligations under article +11 of the WIPO copyright treaty adopted on 20 December 1996, or +similar laws prohibiting or restricting circumvention of such +measures. + + When you convey a covered work, you waive any legal power to forbid +circumvention of technological measures to the extent such circumvention +is effected by exercising rights under this License with respect to +the covered work, and you disclaim any intention to limit operation or +modification of the work as a means of enforcing, against the work's +users, your or third parties' legal rights to forbid circumvention of +technological measures. + + 4. Conveying Verbatim Copies. + + You may convey verbatim copies of the Program's source code as you +receive it, in any medium, provided that you conspicuously and +appropriately publish on each copy an appropriate copyright notice; +keep intact all notices stating that this License and any +non-permissive terms added in accord with section 7 apply to the code; +keep intact all notices of the absence of any warranty; and give all +recipients a copy of this License along with the Program. + + You may charge any price or no price for each copy that you convey, +and you may offer support or warranty protection for a fee. + + 5. Conveying Modified Source Versions. + + You may convey a work based on the Program, or the modifications to +produce it from the Program, in the form of source code under the +terms of section 4, provided that you also meet all of these conditions: + + a) The work must carry prominent notices stating that you modified + it, and giving a relevant date. + + b) The work must carry prominent notices stating that it is + released under this License and any conditions added under section + 7. This requirement modifies the requirement in section 4 to + "keep intact all notices". + + c) You must license the entire work, as a whole, under this + License to anyone who comes into possession of a copy. This + License will therefore apply, along with any applicable section 7 + additional terms, to the whole of the work, and all its parts, + regardless of how they are packaged. This License gives no + permission to license the work in any other way, but it does not + invalidate such permission if you have separately received it. + + d) If the work has interactive user interfaces, each must display + Appropriate Legal Notices; however, if the Program has interactive + interfaces that do not display Appropriate Legal Notices, your + work need not make them do so. + + A compilation of a covered work with other separate and independent +works, which are not by their nature extensions of the covered work, +and which are not combined with it such as to form a larger program, +in or on a volume of a storage or distribution medium, is called an +"aggregate" if the compilation and its resulting copyright are not +used to limit the access or legal rights of the compilation's users +beyond what the individual works permit. Inclusion of a covered work +in an aggregate does not cause this License to apply to the other +parts of the aggregate. + + 6. Conveying Non-Source Forms. + + You may convey a covered work in object code form under the terms +of sections 4 and 5, provided that you also convey the +machine-readable Corresponding Source under the terms of this License, +in one of these ways: + + a) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by the + Corresponding Source fixed on a durable physical medium + customarily used for software interchange. + + b) Convey the object code in, or embodied in, a physical product + (including a physical distribution medium), accompanied by a + written offer, valid for at least three years and valid for as + long as you offer spare parts or customer support for that product + model, to give anyone who possesses the object code either (1) a + copy of the Corresponding Source for all the software in the + product that is covered by this License, on a durable physical + medium customarily used for software interchange, for a price no + more than your reasonable cost of physically performing this + conveying of source, or (2) access to copy the + Corresponding Source from a network server at no charge. + + c) Convey individual copies of the object code with a copy of the + written offer to provide the Corresponding Source. This + alternative is allowed only occasionally and noncommercially, and + only if you received the object code with such an offer, in accord + with subsection 6b. + + d) Convey the object code by offering access from a designated + place (gratis or for a charge), and offer equivalent access to the + Corresponding Source in the same way through the same place at no + further charge. You need not require recipients to copy the + Corresponding Source along with the object code. If the place to + copy the object code is a network server, the Corresponding Source + may be on a different server (operated by you or a third party) + that supports equivalent copying facilities, provided you maintain + clear directions next to the object code saying where to find the + Corresponding Source. Regardless of what server hosts the + Corresponding Source, you remain obligated to ensure that it is + available for as long as needed to satisfy these requirements. + + e) Convey the object code using peer-to-peer transmission, provided + you inform other peers where the object code and Corresponding + Source of the work are being offered to the general public at no + charge under subsection 6d. + + A separable portion of the object code, whose source code is excluded +from the Corresponding Source as a System Library, need not be +included in conveying the object code work. + + A "User Product" is either (1) a "consumer product", which means any +tangible personal property which is normally used for personal, family, +or household purposes, or (2) anything designed or sold for incorporation +into a dwelling. In determining whether a product is a consumer product, +doubtful cases shall be resolved in favor of coverage. For a particular +product received by a particular user, "normally used" refers to a +typical or common use of that class of product, regardless of the status +of the particular user or of the way in which the particular user +actually uses, or expects or is expected to use, the product. A product +is a consumer product regardless of whether the product has substantial +commercial, industrial or non-consumer uses, unless such uses represent +the only significant mode of use of the product. + + "Installation Information" for a User Product means any methods, +procedures, authorization keys, or other information required to install +and execute modified versions of a covered work in that User Product from +a modified version of its Corresponding Source. The information must +suffice to ensure that the continued functioning of the modified object +code is in no case prevented or interfered with solely because +modification has been made. + + If you convey an object code work under this section in, or with, or +specifically for use in, a User Product, and the conveying occurs as +part of a transaction in which the right of possession and use of the +User Product is transferred to the recipient in perpetuity or for a +fixed term (regardless of how the transaction is characterized), the +Corresponding Source conveyed under this section must be accompanied +by the Installation Information. But this requirement does not apply +if neither you nor any third party retains the ability to install +modified object code on the User Product (for example, the work has +been installed in ROM). + + The requirement to provide Installation Information does not include a +requirement to continue to provide support service, warranty, or updates +for a work that has been modified or installed by the recipient, or for +the User Product in which it has been modified or installed. Access to a +network may be denied when the modification itself materially and +adversely affects the operation of the network or violates the rules and +protocols for communication across the network. + + Corresponding Source conveyed, and Installation Information provided, +in accord with this section must be in a format that is publicly +documented (and with an implementation available to the public in +source code form), and must require no special password or key for +unpacking, reading or copying. + + 7. Additional Terms. + + "Additional permissions" are terms that supplement the terms of this +License by making exceptions from one or more of its conditions. +Additional permissions that are applicable to the entire Program shall +be treated as though they were included in this License, to the extent +that they are valid under applicable law. If additional permissions +apply only to part of the Program, that part may be used separately +under those permissions, but the entire Program remains governed by +this License without regard to the additional permissions. + + When you convey a copy of a covered work, you may at your option +remove any additional permissions from that copy, or from any part of +it. (Additional permissions may be written to require their own +removal in certain cases when you modify the work.) You may place +additional permissions on material, added by you to a covered work, +for which you have or can give appropriate copyright permission. + + Notwithstanding any other provision of this License, for material you +add to a covered work, you may (if authorized by the copyright holders of +that material) supplement the terms of this License with terms: + + a) Disclaiming warranty or limiting liability differently from the + terms of sections 15 and 16 of this License; or + + b) Requiring preservation of specified reasonable legal notices or + author attributions in that material or in the Appropriate Legal + Notices displayed by works containing it; or + + c) Prohibiting misrepresentation of the origin of that material, or + requiring that modified versions of such material be marked in + reasonable ways as different from the original version; or + + d) Limiting the use for publicity purposes of names of licensors or + authors of the material; or + + e) Declining to grant rights under trademark law for use of some + trade names, trademarks, or service marks; or + + f) Requiring indemnification of licensors and authors of that + material by anyone who conveys the material (or modified versions of + it) with contractual assumptions of liability to the recipient, for + any liability that these contractual assumptions directly impose on + those licensors and authors. + + All other non-permissive additional terms are considered "further +restrictions" within the meaning of section 10. If the Program as you +received it, or any part of it, contains a notice stating that it is +governed by this License along with a term that is a further +restriction, you may remove that term. If a license document contains +a further restriction but permits relicensing or conveying under this +License, you may add to a covered work material governed by the terms +of that license document, provided that the further restriction does +not survive such relicensing or conveying. + + If you add terms to a covered work in accord with this section, you +must place, in the relevant source files, a statement of the +additional terms that apply to those files, or a notice indicating +where to find the applicable terms. + + Additional terms, permissive or non-permissive, may be stated in the +form of a separately written license, or stated as exceptions; +the above requirements apply either way. + + 8. Termination. + + You may not propagate or modify a covered work except as expressly +provided under this License. Any attempt otherwise to propagate or +modify it is void, and will automatically terminate your rights under +this License (including any patent licenses granted under the third +paragraph of section 11). + + However, if you cease all violation of this License, then your +license from a particular copyright holder is reinstated (a) +provisionally, unless and until the copyright holder explicitly and +finally terminates your license, and (b) permanently, if the copyright +holder fails to notify you of the violation by some reasonable means +prior to 60 days after the cessation. + + Moreover, your license from a particular copyright holder is +reinstated permanently if the copyright holder notifies you of the +violation by some reasonable means, this is the first time you have +received notice of violation of this License (for any work) from that +copyright holder, and you cure the violation prior to 30 days after +your receipt of the notice. + + Termination of your rights under this section does not terminate the +licenses of parties who have received copies or rights from you under +this License. If your rights have been terminated and not permanently +reinstated, you do not qualify to receive new licenses for the same +material under section 10. + + 9. Acceptance Not Required for Having Copies. + + You are not required to accept this License in order to receive or +run a copy of the Program. Ancillary propagation of a covered work +occurring solely as a consequence of using peer-to-peer transmission +to receive a copy likewise does not require acceptance. However, +nothing other than this License grants you permission to propagate or +modify any covered work. These actions infringe copyright if you do +not accept this License. Therefore, by modifying or propagating a +covered work, you indicate your acceptance of this License to do so. + + 10. Automatic Licensing of Downstream Recipients. + + Each time you convey a covered work, the recipient automatically +receives a license from the original licensors, to run, modify and +propagate that work, subject to this License. You are not responsible +for enforcing compliance by third parties with this License. + + An "entity transaction" is a transaction transferring control of an +organization, or substantially all assets of one, or subdividing an +organization, or merging organizations. If propagation of a covered +work results from an entity transaction, each party to that +transaction who receives a copy of the work also receives whatever +licenses to the work the party's predecessor in interest had or could +give under the previous paragraph, plus a right to possession of the +Corresponding Source of the work from the predecessor in interest, if +the predecessor has it or can get it with reasonable efforts. + + You may not impose any further restrictions on the exercise of the +rights granted or affirmed under this License. For example, you may +not impose a license fee, royalty, or other charge for exercise of +rights granted under this License, and you may not initiate litigation +(including a cross-claim or counterclaim in a lawsuit) alleging that +any patent claim is infringed by making, using, selling, offering for +sale, or importing the Program or any portion of it. + + 11. Patents. + + A "contributor" is a copyright holder who authorizes use under this +License of the Program or a work on which the Program is based. The +work thus licensed is called the contributor's "contributor version". + + A contributor's "essential patent claims" are all patent claims +owned or controlled by the contributor, whether already acquired or +hereafter acquired, that would be infringed by some manner, permitted +by this License, of making, using, or selling its contributor version, +but do not include claims that would be infringed only as a +consequence of further modification of the contributor version. For +purposes of this definition, "control" includes the right to grant +patent sublicenses in a manner consistent with the requirements of +this License. + + Each contributor grants you a non-exclusive, worldwide, royalty-free +patent license under the contributor's essential patent claims, to +make, use, sell, offer for sale, import and otherwise run, modify and +propagate the contents of its contributor version. + + In the following three paragraphs, a "patent license" is any express +agreement or commitment, however denominated, not to enforce a patent +(such as an express permission to practice a patent or covenant not to +sue for patent infringement). To "grant" such a patent license to a +party means to make such an agreement or commitment not to enforce a +patent against the party. + + If you convey a covered work, knowingly relying on a patent license, +and the Corresponding Source of the work is not available for anyone +to copy, free of charge and under the terms of this License, through a +publicly available network server or other readily accessible means, +then you must either (1) cause the Corresponding Source to be so +available, or (2) arrange to deprive yourself of the benefit of the +patent license for this particular work, or (3) arrange, in a manner +consistent with the requirements of this License, to extend the patent +license to downstream recipients. "Knowingly relying" means you have +actual knowledge that, but for the patent license, your conveying the +covered work in a country, or your recipient's use of the covered work +in a country, would infringe one or more identifiable patents in that +country that you have reason to believe are valid. + + If, pursuant to or in connection with a single transaction or +arrangement, you convey, or propagate by procuring conveyance of, a +covered work, and grant a patent license to some of the parties +receiving the covered work authorizing them to use, propagate, modify +or convey a specific copy of the covered work, then the patent license +you grant is automatically extended to all recipients of the covered +work and works based on it. + + A patent license is "discriminatory" if it does not include within +the scope of its coverage, prohibits the exercise of, or is +conditioned on the non-exercise of one or more of the rights that are +specifically granted under this License. You may not convey a covered +work if you are a party to an arrangement with a third party that is +in the business of distributing software, under which you make payment +to the third party based on the extent of your activity of conveying +the work, and under which the third party grants, to any of the +parties who would receive the covered work from you, a discriminatory +patent license (a) in connection with copies of the covered work +conveyed by you (or copies made from those copies), or (b) primarily +for and in connection with specific products or compilations that +contain the covered work, unless you entered into that arrangement, +or that patent license was granted, prior to 28 March 2007. + + Nothing in this License shall be construed as excluding or limiting +any implied license or other defenses to infringement that may +otherwise be available to you under applicable patent law. + + 12. No Surrender of Others' Freedom. + + If conditions are imposed on you (whether by court order, agreement or +otherwise) that contradict the conditions of this License, they do not +excuse you from the conditions of this License. If you cannot convey a +covered work so as to satisfy simultaneously your obligations under this +License and any other pertinent obligations, then as a consequence you may +not convey it at all. For example, if you agree to terms that obligate you +to collect a royalty for further conveying from those to whom you convey +the Program, the only way you could satisfy both those terms and this +License would be to refrain entirely from conveying the Program. + + 13. Use with the GNU Affero General Public License. + + Notwithstanding any other provision of this License, you have +permission to link or combine any covered work with a work licensed +under version 3 of the GNU Affero General Public License into a single +combined work, and to convey the resulting work. The terms of this +License will continue to apply to the part which is the covered work, +but the special requirements of the GNU Affero General Public License, +section 13, concerning interaction through a network will apply to the +combination as such. + + 14. Revised Versions of this License. + + The Free Software Foundation may publish revised and/or new versions of +the GNU General Public License from time to time. Such new versions will +be similar in spirit to the present version, but may differ in detail to +address new problems or concerns. + + Each version is given a distinguishing version number. If the +Program specifies that a certain numbered version of the GNU General +Public License "or any later version" applies to it, you have the +option of following the terms and conditions either of that numbered +version or of any later version published by the Free Software +Foundation. If the Program does not specify a version number of the +GNU General Public License, you may choose any version ever published +by the Free Software Foundation. + + If the Program specifies that a proxy can decide which future +versions of the GNU General Public License can be used, that proxy's +public statement of acceptance of a version permanently authorizes you +to choose that version for the Program. + + Later license versions may give you additional or different +permissions. However, no additional obligations are imposed on any +author or copyright holder as a result of your choosing to follow a +later version. + + 15. Disclaimer of Warranty. + + THERE IS NO WARRANTY FOR THE PROGRAM, TO THE EXTENT PERMITTED BY +APPLICABLE LAW. EXCEPT WHEN OTHERWISE STATED IN WRITING THE COPYRIGHT +HOLDERS AND/OR OTHER PARTIES PROVIDE THE PROGRAM "AS IS" WITHOUT WARRANTY +OF ANY KIND, EITHER EXPRESSED OR IMPLIED, INCLUDING, BUT NOT LIMITED TO, +THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR +PURPOSE. THE ENTIRE RISK AS TO THE QUALITY AND PERFORMANCE OF THE PROGRAM +IS WITH YOU. SHOULD THE PROGRAM PROVE DEFECTIVE, YOU ASSUME THE COST OF +ALL NECESSARY SERVICING, REPAIR OR CORRECTION. + + 16. Limitation of Liability. + + IN NO EVENT UNLESS REQUIRED BY APPLICABLE LAW OR AGREED TO IN WRITING +WILL ANY COPYRIGHT HOLDER, OR ANY OTHER PARTY WHO MODIFIES AND/OR CONVEYS +THE PROGRAM AS PERMITTED ABOVE, BE LIABLE TO YOU FOR DAMAGES, INCLUDING ANY +GENERAL, SPECIAL, INCIDENTAL OR CONSEQUENTIAL DAMAGES ARISING OUT OF THE +USE OR INABILITY TO USE THE PROGRAM (INCLUDING BUT NOT LIMITED TO LOSS OF +DATA OR DATA BEING RENDERED INACCURATE OR LOSSES SUSTAINED BY YOU OR THIRD +PARTIES OR A FAILURE OF THE PROGRAM TO OPERATE WITH ANY OTHER PROGRAMS), +EVEN IF SUCH HOLDER OR OTHER PARTY HAS BEEN ADVISED OF THE POSSIBILITY OF +SUCH DAMAGES. + + 17. Interpretation of Sections 15 and 16. + + If the disclaimer of warranty and limitation of liability provided +above cannot be given local legal effect according to their terms, +reviewing courts shall apply local law that most closely approximates +an absolute waiver of all civil liability in connection with the +Program, unless a warranty or assumption of liability accompanies a +copy of the Program in return for a fee. + + END OF TERMS AND CONDITIONS + + How to Apply These Terms to Your New Programs + + If you develop a new program, and you want it to be of the greatest +possible use to the public, the best way to achieve this is to make it +free software which everyone can redistribute and change under these terms. + + To do so, attach the following notices to the program. It is safest +to attach them to the start of each source file to most effectively +state the exclusion of warranty; and each file should have at least +the "copyright" line and a pointer to where the full notice is found. + + + Copyright (C) + + This program is free software: you can redistribute it and/or modify + it under the terms of the GNU General Public License as published by + the Free Software Foundation, either version 3 of the License, or + (at your option) any later version. + + This program is distributed in the hope that it will be useful, + but WITHOUT ANY WARRANTY; without even the implied warranty of + MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the + GNU General Public License for more details. + + You should have received a copy of the GNU General Public License + along with this program. If not, see . + +Also add information on how to contact you by electronic and paper mail. + + If the program does terminal interaction, make it output a short +notice like this when it starts in an interactive mode: + + Copyright (C) + This program comes with ABSOLUTELY NO WARRANTY; for details type `show w'. + This is free software, and you are welcome to redistribute it + under certain conditions; type `show c' for details. + +The hypothetical commands `show w' and `show c' should show the appropriate +parts of the General Public License. Of course, your program's commands +might be different; for a GUI interface, you would use an "about box". + + You should also get your employer (if you work as a programmer) or school, +if any, to sign a "copyright disclaimer" for the program, if necessary. +For more information on this, and how to apply and follow the GNU GPL, see +. + + The GNU General Public License does not permit incorporating your program +into proprietary programs. If your program is a subroutine library, you +may consider it more useful to permit linking proprietary applications with +the library. If this is what you want to do, use the GNU Lesser General +Public License instead of this License. But first, please read +. \ No newline at end of file diff --git a/README.md b/README.md new file mode 100644 index 0000000..c12045e --- /dev/null +++ b/README.md @@ -0,0 +1,119 @@ +# grdp — 纯 Go 的 RDP 客户端协议库 + +本目录是可独立引用的 Go module `git.zeroonesoft.cn/golib/rdplib`(`github.com/nakagami/grdp` 的深度 fork, +MIT 许可,保留上游版权),在 [nakagami/grdp](https://github.com/nakagami/grdp) +(其 fork 自 [tomatome/grdp](https://github.com/tomatome/grdp))基础上大幅扩展, +是 [webrdp](../README.md) 浏览器客户端与原生诊断工具共用的协议栈。 + +## 能力一览(fork 新增/重写部分加粗) + +- RDP 6.0+ 连接:**NLA/CredSSP(NTLMv2)**、TLS、标准安全协商 +- 图形: + - 传统位图/orders 管线(**位图缓存 V2/Memblt/指针形状**) + - **RDPGFX**(ClearCodec/RFX Progressive/Planar/AVC420/AVC444), + **H.264 以原始 NAL 或 I420/NV12 平面回调**交付(对接 WebCodecs/FFmpeg 均可) + - **动态分辨率**(DisplayControl 通道)与**会话色深**切换 +- 输入:扫描码键盘、**Unicode 文本输入**(中文 IME)、鼠标移动合并 +- 剪贴板:文本/图片/HTML/**文件复制**(含进度回调) +- 音频:rdpsnd + AUDIO_PLAYBACK DVC,三模式(本地/远端/静音) +- **驱动器重定向(rdpdr)**:宣告/验证/IO 编码、只读文件系统接口 + (状态见 ../doc/RDPDR-2.md) +- 连接期自动检测(RTT/带宽测量应答)、**并行 MCS 通道加入**、 + 服务器重定向 PDU 处理 +- **事件发射器 emission**:反射最小化 + 快速路径可摘除(有回归测试) +- 诊断:PDU 录制器、GFX 缓存存储接口 + +## 快速开始(原生 Go 调用) + +```go +package main + +import ( + "fmt" + "net" + + "github.com/nakagami/grdp" +) + +func main() { + client := grdp.NewRdpClient("10.0.0.3:3389", 1280, 800, + func(addr string) (net.Conn, error) { return net.Dial("tcp", addr) }) + + client.OnError(func(e error) { fmt.Println("err:", e) }). + OnClose(func() { fmt.Println("closed") }). + OnReady(func() { fmt.Println("session ready") }). + OnBitmap(func(bits []grdp.Bitmap) { /* RGBA() / FillRGBA(dst) */ }). + OnH264I420(func(x, y, w, h int, yb []byte, ys, u []byte, us, v []byte, vs int) { + // 送入渲染器/编码器 + }). + OnClipboard(func(text string) { /* 远端剪贴板文本 */ }, func() string { return "" }) + + if err := client.Login("", "administrator", "password"); err != nil { + panic(err) + } + // 输入/分辨率/关闭: + // client.SendUnicodeText("你好"); client.SetResolution(1920, 1080); client.Close() + select {} +} +``` + +完整可运行示例:仓库上层 `cmd/nativedemo`(带带宽/编码统计的诊断客户端, +`go run ./cmd/nativedemo -host x -user u -pass p`)。 +浏览器(WASM/WebSocket)宿主实现:仓库上层根目录 `main.go` + `wasm_transport.go`。 + +回调都在协议栈读循环 goroutine 内触发——回调里不要长时间阻塞, +重活请转交自己的 goroutine。 + +## 架构 + +``` +grdp.RdpClient 门户:配置(链式 Set*/On*) + Login + 生命周期 +├─ protocol/x224·tpkt·t125·sec 传输/安全层(MCS 通道、加密;NLA 在 protocol/nla) +├─ protocol/pdu·lic·gcc 虚拟桌面 PDU、能力协商、许可、GCC +├─ plugin/ 虚拟通道框架 + 内置通道 +│ ├─ Channels.Register(ChannelTransport) 通道注册 +│ └─ rdpsnd cliprdr rdpdr rdpgfx rdpedisp drdynvc +├─ core/ 传输抽象、缓冲池、工具 +└─ emission/ 事件发射器 +``` + +自定义虚拟通道只需实现三个方法(见 `plugin/channel.go`): + +```go +type ChannelTransport interface { + GetType() (string, uint32) // 通道名 + CHANNEL_OPTION_* + Sender(core.ChannelSender) // 栈回调的发送器 + Process(s []byte) // 收包 +} +``` + +## 构建与测试 + +``` +go build ./... # 普通 GOOS 即可;WASM 用 GOOS=js GOARCH=wasm +go test ./... +``` + +FFmpeg 硬解(可选):`-tags h264`,依赖 libavcodec ≥3.4(macOS VideoToolbox / +Linux VAAPI 自动启用,软件回退)。纯 Go 构建默认不带 AVC 解码 +(但保留 AVC 能力协商与 NAL 回调,解码交给宿主,如浏览器的 WebCodecs)。 + +## 发布检查单(fork 注意事项) + +- 模块路径当前沿用上游 `github.com/nakagami/grdp`(可离线构建、自主维护, + 经上层 go.mod `replace` 引用)。**对外发布前需改为你自己的 module path**: + 改本目录 `go.mod` 的 module 行 + 全部 import 前缀 + (`sed -i 's|github.com/nakagami/grdp|<你的路径>|g'`),与上层仓库无其他耦合。 +- 许可证 MIT,上游版权声明见 LICENSE,请保留。 +- 已知问题:驱动器重定向在部分 Windows 服务端上被关闭(RDPDR-2), + 跟踪见 ../doc/RDPDR-2.md。 + +## 上游致谢与相关项目 + +- 上游取材:[rdpy](https://github.com/citronneur/rdpy)、 + [node-rdpjs](https://github.com/citronneur/node-rdpjs)、 + [gordp](https://github.com/Madnikulin50/gordp)、 + [ncrack_rdp](https://github.com/nmap/ncrack/blob/master/modules/ncrack_rdp.cc)、 + [webRDP](https://github.com/Chorder/webRDP) +- 相关:https://github.com/nakagami/grdpsdl2 、 + https://github.com/nakagami/grdpwasm diff --git a/bmpcache_test.go b/bmpcache_test.go new file mode 100644 index 0000000..7086296 --- /dev/null +++ b/bmpcache_test.go @@ -0,0 +1,105 @@ +package grdp + +import ( + "testing" + + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpgfx" + "git.zeroonesoft.cn/golib/rdplib/protocol/pdu" +) + +// fakeBmpStore 记录 Persist 调用并提供 Get/Keys,用于验证 +// CacheBitmapV2 持久键的跨会话存取(6.4b M2)。 +type fakeBmpStore struct { + persisted map[uint64]rdpgfx.GfxCacheEntry +} + +func newFakeBmpStore() *fakeBmpStore { + return &fakeBmpStore{persisted: map[uint64]rdpgfx.GfxCacheEntry{}} +} + +func (f *fakeBmpStore) Persist(key uint64, w, h int, bpp uint16, data []byte) { + f.persisted[key] = rdpgfx.GfxCacheEntry{Key: key, Width: w, Height: h, Bpp: bpp, Data: data} +} +func (f *fakeBmpStore) Export() []rdpgfx.GfxCacheEntry { + out := make([]rdpgfx.GfxCacheEntry, 0, len(f.persisted)) + for _, e := range f.persisted { + out = append(out, e) + } + return out +} +func (f *fakeBmpStore) Get(key uint64) (rdpgfx.GfxCacheEntry, bool) { + e, ok := f.persisted[key] + return e, ok +} +func (f *fakeBmpStore) Keys() []uint64 { + out := make([]uint64, 0, len(f.persisted)) + for k := range f.persisted { + out = append(out, k) + } + return out +} + +// TestStoreBitmapCacheV2PersistentKey:带持久键的 CacheBitmapV2 应把 +// 位图以 key2<<32|key1 持久化;零数据+持久键应从持久库回填会话内单元。 +func TestStoreBitmapCacheV2PersistentKey(t *testing.T) { + store := newFakeBmpStore() + g := &RdpClient{ + bitmapCache: make(map[uint32]*Bitmap), + bitmapCacheFIFO: []uint32{}, + gfxCacheStore: store, + } + + // 1) 带键+数据:进会话缓存并持久化 + cb := &pdu.CacheBitmapV2Order{ + CacheId: 1, + Flags: pdu.CBR2_PERSISTENT_KEY_PRESENT, + Key1: 0xDEAD, + Key2: 0xBEEF, + BitmapBpp: 32, + BitmapWidth: 4, + BitmapHeight: 4, + BitmapLength: 64, + CacheIndex: 7, + BitmapDataStream: make([]byte, 4*4*4), + } + g.storeBitmapCacheV2(cb) + pk := uint64(0xBEEF)<<32 | 0xDEAD + if _, ok := store.persisted[pk]; !ok { + t.Fatalf("entry not persisted under %x; got %v", pk, store.persisted) + } + if _, ok := g.bitmapCache[1<<16|7]; !ok { + t.Fatal("entry missing from in-session cache") + } + + // 2) 零数据+同键(另一单元):从持久库回填 + cb2 := &pdu.CacheBitmapV2Order{ + CacheId: 2, + Flags: pdu.CBR2_PERSISTENT_KEY_PRESENT, + Key1: 0xDEAD, + Key2: 0xBEEF, + BitmapBpp: 32, + BitmapWidth: 4, + BitmapHeight: 4, + CacheIndex: 9, + } + g.storeBitmapCacheV2(cb2) + b, ok := g.bitmapCache[2<<16|9] + if !ok || b == nil || len(b.Data) != 4*4*4 { + t.Fatalf("persistent backfill failed: ok=%v len=%v", ok, b) + } + + // 3) 零数据但持久库未命中:不放入缓存 + cb3 := &pdu.CacheBitmapV2Order{ + CacheId: 3, + Flags: pdu.CBR2_PERSISTENT_KEY_PRESENT, + Key1: 0x1111, + Key2: 0x2222, + BitmapBpp: 32, + BitmapWidth: 4, BitmapHeight: 4, + CacheIndex: 5, + } + g.storeBitmapCacheV2(cb3) + if _, ok := g.bitmapCache[3<<16|5]; ok { + t.Fatal("missed backfill must not populate cache") + } +} diff --git a/convert_amd64.go b/convert_amd64.go new file mode 100644 index 0000000..2e2888d --- /dev/null +++ b/convert_amd64.go @@ -0,0 +1,55 @@ +//go:build amd64 + +package grdp + +import "encoding/binary" + +// bgr32BatchToRGBA converts n BGRA32 pixels (4 bytes each: B,G,R,X in memory) to +// RGBA using SSE2. Processes 8 pixels per iteration; remainder is scalar. +func bgr32BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + bgr32toRGBAasm(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + s := i * 4 + binary.LittleEndian.PutUint32(dst[s:], + uint32(src[s+2])|uint32(src[s+1])<<8|uint32(src[s])<<16|0xFF000000) + } +} + +//go:noescape +func bgr32toRGBAasm(dst *byte, src *byte, n int) + +// rgb555BatchToRGBA converts n big-endian RGB555 pixels to RGBA using SSE2. +// Processes 8 pixels per iteration; any remainder is handled via scalar fallback. +func rgb555BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + rgb555toRGBAasm(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0x7C00)>>7)|uint32((d&0x03E0)>>2)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} + +// rgb565BatchToRGBA converts n big-endian RGB565 pixels to RGBA using SSE2. +func rgb565BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + rgb565toRGBAasm(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0xF800)>>8)|uint32((d&0x07E0)>>3)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} + +//go:noescape +func rgb555toRGBAasm(dst *byte, src *byte, n int) + +//go:noescape +func rgb565toRGBAasm(dst *byte, src *byte, n int) diff --git a/convert_amd64.s b/convert_amd64.s new file mode 100644 index 0000000..6e747ca --- /dev/null +++ b/convert_amd64.s @@ -0,0 +1,210 @@ +// SSE2 implementations of bgr32toRGBAasm, rgb555toRGBAasm, and rgb565toRGBAasm. +// Each function processes 8 big-endian RGB pixels per loop iteration. +// Stack ABI (ABI0): dst+0(FP), src+8(FP), n+16(FP) — total 24 bytes. + +#include "textflag.h" + +// Packed-word masks used across both functions. +DATA rgb_00F8<>+0x00(SB)/8, $0x00F800F800F800F8 +DATA rgb_00F8<>+0x08(SB)/8, $0x00F800F800F800F8 +GLOBL rgb_00F8<>(SB), (NOPTR|RODATA), $16 + +DATA rgb_00FC<>+0x00(SB)/8, $0x00FC00FC00FC00FC +DATA rgb_00FC<>+0x08(SB)/8, $0x00FC00FC00FC00FC +GLOBL rgb_00FC<>(SB), (NOPTR|RODATA), $16 + +DATA rgb_FF00<>+0x00(SB)/8, $0xFF00FF00FF00FF00 +DATA rgb_FF00<>+0x08(SB)/8, $0xFF00FF00FF00FF00 +GLOBL rgb_FF00<>(SB), (NOPTR|RODATA), $16 + +// func rgb555toRGBAasm(dst *byte, src *byte, n int) +// Converts n big-endian RGB555 pixels to RGBA. +// R = (d>>7)&0xF8, G = (d>>2)&0xF8, B = (d<<3)&0xF8, A = 0xFF. +TEXT ·rgb555toRGBAasm(SB),NOSPLIT,$0-24 + MOVQ dst+0(FP), DI + MOVQ src+8(FP), SI + MOVQ n+16(FP), AX + MOVOU rgb_00F8<>(SB), X13 // mask 0x00F8 in each 16-bit lane + MOVOU rgb_FF00<>(SB), X14 // mask 0xFF00 in each 16-bit lane + +loop555: + // Load 16 bytes = 8 big-endian uint16 pixels. + MOVOU (SI), X0 + + // Byte-swap each 16-bit element: memory has [H,L] per pixel; + // x86 loads give word = L<<8|H; we want d = H<<8|L. + MOVO X0, X1 + PSLLW $8, X0 // X0[i] = H<<8 (low byte cleared) + PSRLW $8, X1 // X1[i] = L (high byte cleared) + POR X1, X0 // X0[i] = H<<8|L = d + + // Extract R = (d>>7) & 0x00F8. + MOVO X0, X2 + PSRLW $7, X2 + PAND X13, X2 + + // Extract G = (d>>2) & 0x00F8. + MOVO X0, X3 + PSRLW $2, X3 + PAND X13, X3 + + // Extract B = (d<<3) & 0x00F8. + MOVO X0, X4 + PSLLW $3, X4 + PAND X13, X4 + + // Build RG word: G in high byte, R in low byte. + MOVO X3, X5 + PSLLW $8, X5 // G → high byte + POR X2, X5 // X5[i] = G<<8|R + + // Build BA word: 0xFF in high byte, B in low byte. + MOVO X4, X6 + POR X14, X6 // X6[i] = 0xFF00|B + + // Interleave to produce 4 RGBA dwords each. + // PUNPCKLWL src,dst → dst=[dst_w0,src_w0,dst_w1,src_w1,...dst_w3,src_w3] + // Each dword becomes bytes [R,G,B,0xFF]. + MOVO X5, X7 + PUNPCKLWL X6, X5 // low 4 pixels → X5 + PUNPCKHWL X6, X7 // high 4 pixels → X7 + + MOVOU X5, (DI) + MOVOU X7, 16(DI) + + ADDQ $16, SI + ADDQ $32, DI + SUBQ $8, AX + JNZ loop555 + RET + +// func rgb565toRGBAasm(dst *byte, src *byte, n int) +// Converts n big-endian RGB565 pixels to RGBA. +// R = (d>>8)&0xF8, G = (d>>3)&0xFC, B = (d<<3)&0xF8, A = 0xFF. +TEXT ·rgb565toRGBAasm(SB),NOSPLIT,$0-24 + MOVQ dst+0(FP), DI + MOVQ src+8(FP), SI + MOVQ n+16(FP), AX + MOVOU rgb_00F8<>(SB), X13 + MOVOU rgb_00FC<>(SB), X15 + MOVOU rgb_FF00<>(SB), X14 + +loop565: + MOVOU (SI), X0 + + MOVO X0, X1 + PSLLW $8, X0 + PSRLW $8, X1 + POR X1, X0 // X0[i] = d + + // R = (d>>8) & 0x00F8. + MOVO X0, X2 + PSRLW $8, X2 + PAND X13, X2 + + // G = (d>>3) & 0x00FC. + MOVO X0, X3 + PSRLW $3, X3 + PAND X15, X3 + + // B = (d<<3) & 0x00F8. + MOVO X0, X4 + PSLLW $3, X4 + PAND X13, X4 + + MOVO X3, X5 + PSLLW $8, X5 + POR X2, X5 // X5[i] = G<<8|R + + MOVO X4, X6 + POR X14, X6 // X6[i] = 0xFF00|B + + MOVO X5, X7 + PUNPCKLWL X6, X5 + PUNPCKHWL X6, X7 + + MOVOU X5, (DI) + MOVOU X7, 16(DI) + + ADDQ $16, SI + ADDQ $32, DI + SUBQ $8, AX + JNZ loop565 + RET + +// Dword-lane masks for bgr32toRGBAasm. +// bgr32_lo: 0x000000FF in each 32-bit lane — isolates the low byte (B or R after shift). +DATA bgr32_lo<>+0x00(SB)/8, $0x000000FF000000FF +DATA bgr32_lo<>+0x08(SB)/8, $0x000000FF000000FF +GLOBL bgr32_lo<>(SB), (NOPTR|RODATA), $16 + +// bgr32_gg: 0x0000FF00 in each 32-bit lane — isolates the G byte. +DATA bgr32_gg<>+0x00(SB)/8, $0x0000FF000000FF00 +DATA bgr32_gg<>+0x08(SB)/8, $0x0000FF000000FF00 +GLOBL bgr32_gg<>(SB), (NOPTR|RODATA), $16 + +// bgr32_aa: 0xFF000000 in each 32-bit lane — supplies the alpha byte. +DATA bgr32_aa<>+0x00(SB)/8, $0xFF000000FF000000 +DATA bgr32_aa<>+0x08(SB)/8, $0xFF000000FF000000 +GLOBL bgr32_aa<>(SB), (NOPTR|RODATA), $16 + +// func bgr32toRGBAasm(dst *byte, src *byte, n int) +// Converts n BGRA32 pixels (memory layout B,G,R,X per pixel) to RGBA +// (memory layout R,G,B,0xFF). Processes 8 pixels (32 bytes) per iteration. +// Source dword in register (LE): bits[7:0]=B, bits[15:8]=G, bits[23:16]=R, bits[31:24]=X. +// Dest dword in register (LE): bits[7:0]=R, bits[15:8]=G, bits[23:16]=B, bits[31:24]=FF. +TEXT ·bgr32toRGBAasm(SB),NOSPLIT,$0-24 + MOVQ dst+0(FP), DI + MOVQ src+8(FP), SI + MOVQ n+16(FP), AX + MOVOU bgr32_lo<>(SB), X12 // 0x000000FF per dword + MOVOU bgr32_gg<>(SB), X13 // 0x0000FF00 per dword + MOVOU bgr32_aa<>(SB), X14 // 0xFF000000 per dword + +loop32: + // ---- pixels 0–3 (16 bytes at SI) ---- + MOVOU (SI), X0 + + // R = (src >> 16) & 0xFF → low byte of dest dword. + MOVO X0, X2 + PSRLL $16, X2 + PAND X12, X2 // X2 = [R,0,0,0] per dword + + // G stays at byte 1: src & 0x0000FF00. + MOVO X0, X3 + PAND X13, X3 // X3 = [0,G,0,0] per dword + + // B moves from byte 0 to byte 2: (src & 0xFF) << 16. + MOVO X0, X4 + PAND X12, X4 + PSLLL $16, X4 // X4 = [0,0,B,0] per dword + + POR X3, X2 + POR X4, X2 + POR X14, X2 // X2 = [R,G,B,FF] per dword + MOVOU X2, (DI) + + // ---- pixels 4–7 (next 16 bytes) ---- + MOVOU 16(SI), X0 + + MOVO X0, X2 + PSRLL $16, X2 + PAND X12, X2 + + MOVO X0, X3 + PAND X13, X3 + + MOVO X0, X4 + PAND X12, X4 + PSLLL $16, X4 + + POR X3, X2 + POR X4, X2 + POR X14, X2 + MOVOU X2, 16(DI) + + ADDQ $32, SI + ADDQ $32, DI + SUBQ $8, AX + JNZ loop32 + RET diff --git a/convert_arm64.go b/convert_arm64.go new file mode 100644 index 0000000..a63da05 --- /dev/null +++ b/convert_arm64.go @@ -0,0 +1,52 @@ +//go:build arm64 + +package grdp + +import "encoding/binary" + +// bgr32BatchToRGBA converts n BGRA32 pixels (4 bytes each: B,G,R,X in memory) to +// RGBA using NEON. Processes 8 pixels per iteration; remainder is scalar. +func bgr32BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + bgr32toRGBAarm64(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + s := i * 4 + binary.LittleEndian.PutUint32(dst[s:], + uint32(src[s+2])|uint32(src[s+1])<<8|uint32(src[s])<<16|0xFF000000) + } +} + +//go:noescape +func bgr32toRGBAarm64(dst *byte, src *byte, n int) + +func rgb555BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + rgb555toRGBAarm64(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0x7C00)>>7)|uint32((d&0x03E0)>>2)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} + +func rgb565BatchToRGBA(dst []byte, src []byte, n int) { + n8 := n &^ 7 + if n8 > 0 { + rgb565toRGBAarm64(&dst[0], &src[0], n8) + } + for i := n8; i < n; i++ { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0xF800)>>8)|uint32((d&0x07E0)>>3)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} + +//go:noescape +func rgb555toRGBAarm64(dst *byte, src *byte, n int) + +//go:noescape +func rgb565toRGBAarm64(dst *byte, src *byte, n int) diff --git a/convert_arm64.s b/convert_arm64.s new file mode 100644 index 0000000..a83e841 --- /dev/null +++ b/convert_arm64.s @@ -0,0 +1,171 @@ +// NEON implementations of rgb555toRGBAarm64 and rgb565toRGBAarm64. +// Processes 8 big-endian RGB pixels per loop iteration. +// Stack ABI (ABI0): dst+0(FP), src+8(FP), n+16(FP) — total 24 bytes. +// +// Instruction notes (Go arm64 assembler requires V-prefix for NEON): +// VUSHR/VSHL = unsigned shift right/left by immediate on vector register +// VUZP1 = unzip even elements; used here to narrow H8→B8 (XTN equivalent) +// VZIP1/VZIP2 = interleave lower/upper halves of two vector registers +// VMOVI $n, Vd.B16 = broadcast 8-bit immediate to all 16 byte lanes + +#include "textflag.h" + +// func rgb555toRGBAarm64(dst *byte, src *byte, n int) +TEXT ·rgb555toRGBAarm64(SB),NOSPLIT,$0-24 + MOVD dst+0(FP), R0 + MOVD src+8(FP), R1 + MOVD n+16(FP), R2 + + // V14 = 0xF8 in every byte (5-bit channel mask). + // V13 = 0xFF in every byte (alpha). + VMOVI $0xF8, V14.B16 + VMOVI $0xFF, V13.B16 + +loop555: + // Load 16 bytes (8 big-endian uint16 pixels); post-increment R1 by 16. + VLD1.P 16(R1), [V0.B16] + + // Byte-swap each 16-bit element: big-endian → native uint16. + // Memory: [H0,L0,H1,L1,...]; after VREV16: V0.H[i] = H_i<<8|L_i = d_i. + VREV16 V0.B16, V0.B16 + + // Extract R = (d>>7) & 0xF8: d bits[14:10] → output bits[7:3]. + // VUSHR gives V1.H[i]=d>>7; VAND zeros high byte; VUZP1 narrows to B8. + VUSHR $7, V0.H8, V1.H8 + VAND V14.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V2.B16 // V2.B[0..7] = R0..R7 + + // Extract G = (d>>2) & 0xF8: d bits[9:5] → output bits[7:3]. + VUSHR $2, V0.H8, V1.H8 + VAND V14.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V3.B16 // V3.B[0..7] = G0..G7 + + // Extract B = (d<<3) & 0xF8: d bits[4:0] → output bits[7:3]. + VSHL $3, V0.H8, V1.H8 + VAND V14.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V4.B16 // V4.B[0..7] = B0..B7 + + // Interleave R and G: [R0,G0,R1,G1,...,R7,G7] (uses lower 8 bytes of each). + VZIP1 V2.B16, V3.B16, V5.B16 + + // Interleave B and alpha: [B0,FF,B1,FF,...,B7,FF]. + VZIP1 V4.B16, V13.B16, V6.B16 + + // Interleave RG and BA halfwords to form RGBA dwords. + VZIP1 V5.H8, V6.H8, V7.H8 // V7 = first 4 RGBA pixels + VZIP2 V5.H8, V6.H8, V8.H8 // V8 = last 4 RGBA pixels + + // Store 32 bytes to dst; post-increment R0 by 32. + VST1.P [V7.B16, V8.B16], 32(R0) + + SUBS $8, R2, R2 + BNE loop555 + RET + +// func rgb565toRGBAarm64(dst *byte, src *byte, n int) +TEXT ·rgb565toRGBAarm64(SB),NOSPLIT,$0-24 + MOVD dst+0(FP), R0 + MOVD src+8(FP), R1 + MOVD n+16(FP), R2 + + VMOVI $0xF8, V14.B16 // 5-bit channel mask (R and B) + VMOVI $0xFC, V12.B16 // 6-bit channel mask (G) + VMOVI $0xFF, V13.B16 // alpha + +loop565: + VLD1.P 16(R1), [V0.B16] + VREV16 V0.B16, V0.B16 + + // R = (d>>8) & 0xF8: d bits[15:11] → output bits[7:3]. + VUSHR $8, V0.H8, V1.H8 + VAND V14.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V2.B16 // V2.B[0..7] = R0..R7 + + // G = (d>>3) & 0xFC: d bits[10:5] → output bits[7:2]. + VUSHR $3, V0.H8, V1.H8 + VAND V12.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V3.B16 // V3.B[0..7] = G0..G7 + + // B = (d<<3) & 0xF8: d bits[4:0] → output bits[7:3]. + VSHL $3, V0.H8, V1.H8 + VAND V14.B16, V1.B16, V1.B16 + VUZP1 V1.B16, V1.B16, V4.B16 // V4.B[0..7] = B0..B7 + + VZIP1 V2.B16, V3.B16, V5.B16 + VZIP1 V4.B16, V13.B16, V6.B16 + + VZIP1 V5.H8, V6.H8, V7.H8 + VZIP2 V5.H8, V6.H8, V8.H8 + + VST1.P [V7.B16, V8.B16], 32(R0) + + SUBS $8, R2, R2 + BNE loop565 + RET + +// bgr32_s4_lo: 0x000000FF in each 32-bit lane — isolates the low byte (B or R after shift). +DATA bgr32_s4_lo<>+0x00(SB)/8, $0x000000FF000000FF +DATA bgr32_s4_lo<>+0x08(SB)/8, $0x000000FF000000FF +GLOBL bgr32_s4_lo<>(SB), (NOPTR|RODATA), $16 + +// bgr32_s4_gg: 0x0000FF00 in each 32-bit lane — isolates the G byte. +DATA bgr32_s4_gg<>+0x00(SB)/8, $0x0000FF000000FF00 +DATA bgr32_s4_gg<>+0x08(SB)/8, $0x0000FF000000FF00 +GLOBL bgr32_s4_gg<>(SB), (NOPTR|RODATA), $16 + +// bgr32_s4_aa: 0xFF000000 in each 32-bit lane — supplies the alpha byte. +DATA bgr32_s4_aa<>+0x00(SB)/8, $0xFF000000FF000000 +DATA bgr32_s4_aa<>+0x08(SB)/8, $0xFF000000FF000000 +GLOBL bgr32_s4_aa<>(SB), (NOPTR|RODATA), $16 + +// func bgr32toRGBAarm64(dst *byte, src *byte, n int) +// Converts n BGRA32 pixels (memory layout B,G,R,X per pixel) to RGBA +// (memory layout R,G,B,0xFF). Processes 8 pixels (32 bytes) per iteration. +// +// Strategy (mirrors amd64 SSE2 implementation — per-dword shift+mask): +// Source dword (LE register): bits[7:0]=B, bits[15:8]=G, bits[23:16]=R, bits[31:24]=X +// Dest dword (LE register): bits[7:0]=R, bits[15:8]=G, bits[23:16]=B, bits[31:24]=FF +// R = (src >> 16) & 0x000000FF +// G = src & 0x0000FF00 +// B = (src & 0x000000FF) << 16 +// A = 0xFF000000 +TEXT ·bgr32toRGBAarm64(SB),NOSPLIT,$0-24 + MOVD dst+0(FP), R0 + MOVD src+8(FP), R1 + MOVD n+16(FP), R2 + + MOVD $bgr32_s4_lo<>(SB), R10 + VLD1 (R10), [V15.B16] // V15 = 0x000000FF per dword + MOVD $bgr32_s4_gg<>(SB), R10 + VLD1 (R10), [V14.B16] // V14 = 0x0000FF00 per dword + MOVD $bgr32_s4_aa<>(SB), R10 + VLD1 (R10), [V13.B16] // V13 = 0xFF000000 per dword + +loop32: + // pixels 0-3 + VLD1.P 16(R1), [V0.B16] + VUSHR $16, V0.S4, V1.S4 // V1 = src >> 16 per dword + VAND V15.B16, V1.B16, V1.B16 // V1 = [R,0,0,0] per dword + VAND V14.B16, V0.B16, V2.B16 // V2 = [0,G,0,0] per dword + VAND V15.B16, V0.B16, V3.B16 // V3 = [B,0,0,0] per dword + VSHL $16, V3.S4, V3.S4 // V3 = [0,0,B,0] per dword + VORR V2.B16, V1.B16, V1.B16 + VORR V3.B16, V1.B16, V1.B16 + VORR V13.B16, V1.B16, V1.B16 // V1 = [R,G,B,FF] per dword + VST1.P [V1.B16], 16(R0) + + // pixels 4-7 + VLD1.P 16(R1), [V0.B16] + VUSHR $16, V0.S4, V1.S4 + VAND V15.B16, V1.B16, V1.B16 + VAND V14.B16, V0.B16, V2.B16 + VAND V15.B16, V0.B16, V3.B16 + VSHL $16, V3.S4, V3.S4 + VORR V2.B16, V1.B16, V1.B16 + VORR V3.B16, V1.B16, V1.B16 + VORR V13.B16, V1.B16, V1.B16 + VST1.P [V1.B16], 16(R0) + + SUBS $8, R2, R2 + BNE loop32 + RET diff --git a/convert_generic.go b/convert_generic.go new file mode 100644 index 0000000..d285829 --- /dev/null +++ b/convert_generic.go @@ -0,0 +1,33 @@ +//go:build !amd64 && !arm64 + +package grdp + +import "encoding/binary" + +// bgr32BatchToRGBA converts n BGRA32 pixels (4 bytes each: B,G,R,X) to RGBA. +func bgr32BatchToRGBA(dst []byte, src []byte, n int) { + for i := range n { + s := i * 4 + binary.LittleEndian.PutUint32(dst[s:], + uint32(src[s+2])|uint32(src[s+1])<<8|uint32(src[s])<<16|0xFF000000) + } +} + +// rgb555BatchToRGBA converts n big-endian RGB555 pixels (src, 2 bytes each) +// to RGBA (dst, 4 bytes each). n must be valid for the slice sizes. +func rgb555BatchToRGBA(dst []byte, src []byte, n int) { + for i := range n { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0x7C00)>>7)|uint32((d&0x03E0)>>2)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} + +// rgb565BatchToRGBA converts n big-endian RGB565 pixels to RGBA. +func rgb565BatchToRGBA(dst []byte, src []byte, n int) { + for i := range n { + d := binary.BigEndian.Uint16(src[i*2:]) + binary.LittleEndian.PutUint32(dst[i*4:], + uint32((d&0xF800)>>8)|uint32((d&0x07E0)>>3)<<8|uint32((d&0x001F)<<3)<<16|0xFF000000) + } +} diff --git a/core/io.go b/core/io.go new file mode 100644 index 0000000..bf1aa20 --- /dev/null +++ b/core/io.go @@ -0,0 +1,131 @@ +package core + +import ( + "encoding/binary" + "io" +) + +type ReadBytesComplete func(result []byte, err error) + +func StartReadBytes(len int, r io.Reader, cb ReadBytesComplete) { + b := make([]byte, len) + go func() { + _, err := io.ReadFull(r, b) + cb(b, err) + }() +} + +func ReadBytes(len int, r io.Reader) ([]byte, error) { + b := make([]byte, len) + length, err := io.ReadFull(r, b) + return b[:length], err +} + +func ReadByte(r io.Reader) (byte, error) { + var buf [1]byte + _, err := io.ReadFull(r, buf[:]) + return buf[0], err +} + +func ReadUInt8(r io.Reader) (uint8, error) { + var buf [1]byte + _, err := io.ReadFull(r, buf[:]) + return buf[0], err +} + +func ReadUint16LE(r io.Reader) (uint16, error) { + var buf [2]byte + _, err := io.ReadFull(r, buf[:]) + if err != nil { + return 0, err + } + return binary.LittleEndian.Uint16(buf[:]), nil +} + +func ReadUint16BE(r io.Reader) (uint16, error) { + var buf [2]byte + _, err := io.ReadFull(r, buf[:]) + if err != nil { + return 0, err + } + return binary.BigEndian.Uint16(buf[:]), nil +} + +func ReadUInt32LE(r io.Reader) (uint32, error) { + var buf [4]byte + _, err := io.ReadFull(r, buf[:]) + if err != nil { + return 0, err + } + return binary.LittleEndian.Uint32(buf[:]), nil +} + +func ReadUInt32BE(r io.Reader) (uint32, error) { + var buf [4]byte + _, err := io.ReadFull(r, buf[:]) + if err != nil { + return 0, err + } + return binary.BigEndian.Uint32(buf[:]), nil +} + +func WriteByte(data byte, w io.Writer) (int, error) { + buf := [1]byte{data} + return w.Write(buf[:]) +} + +func WriteBytes(data []byte, w io.Writer) (int, error) { + return w.Write(data) +} + +func WriteUInt8(data uint8, w io.Writer) (int, error) { + buf := [1]byte{data} + return w.Write(buf[:]) +} + +func WriteUInt16BE(data uint16, w io.Writer) (int, error) { + var buf [2]byte + binary.BigEndian.PutUint16(buf[:], data) + return w.Write(buf[:]) +} + +func WriteUInt16LE(data uint16, w io.Writer) (int, error) { + var buf [2]byte + binary.LittleEndian.PutUint16(buf[:], data) + return w.Write(buf[:]) +} + +func WriteUInt32LE(data uint32, w io.Writer) (int, error) { + var buf [4]byte + binary.LittleEndian.PutUint32(buf[:], data) + return w.Write(buf[:]) +} + +func WriteUInt32BE(data uint32, w io.Writer) (int, error) { + var buf [4]byte + binary.BigEndian.PutUint32(buf[:], data) + return w.Write(buf[:]) +} + +func PutUint16BE(data uint16) (uint8, uint8) { + return uint8(data >> 8), uint8(data) +} + +func Uint16BE(d0, d1 uint8) uint16 { + return uint16(d0)<<8 | uint16(d1) +} + +func RGB565ToRGB(data uint16) (r, g, b uint8) { + r = uint8((data & 0xF800) >> 8) + g = uint8((data & 0x07E0) >> 3) + b = uint8((data & 0x001F) << 3) + + return +} +func RGB555ToRGB(data uint16) (r, g, b uint8) { + r = uint8((data & 0x7C00) >> 7) + g = uint8((data & 0x03E0) >> 2) + b = uint8((data & 0x001F) << 3) + + return +} diff --git a/core/io_test.go b/core/io_test.go new file mode 100644 index 0000000..fc2731b --- /dev/null +++ b/core/io_test.go @@ -0,0 +1,17 @@ +package core + +import ( + "bytes" + "encoding/hex" + "testing" +) + +func TestWriteUInt16LE(t *testing.T) { + buff := &bytes.Buffer{} + WriteUInt32LE(66538, buff) + result := hex.EncodeToString(buff.Bytes()) + expected := "ea030100" + if result != expected { + t.Error(result, "not equals to", expected) + } +} diff --git a/core/log.go b/core/log.go new file mode 100644 index 0000000..35dd4de --- /dev/null +++ b/core/log.go @@ -0,0 +1,20 @@ +package core + +import ( + "encoding/hex" + "log/slog" +) + +// Hex wraps a byte slice as a slog.LogValuer that lazily encodes the bytes +// as a hexadecimal string only when the slog handler actually formats it. +// +// Use Hex(buf) instead of hex.EncodeToString(buf) inside slog.Debug calls on +// hot paths: when the logger's level filter discards the record (the common +// case in production), the encode and the per-call string allocation are +// skipped entirely. +type Hex []byte + +// LogValue implements slog.LogValuer. +func (h Hex) LogValue() slog.Value { + return slog.StringValue(hex.EncodeToString(h)) +} diff --git a/core/mppc.go b/core/mppc.go new file mode 100644 index 0000000..c394311 --- /dev/null +++ b/core/mppc.go @@ -0,0 +1,205 @@ +package core + +import ( + "errors" + "fmt" +) + +// MppcDecompressor maintains per-connection state for RDP bulk data +// decompression (MS-RDPBCGR §3.1.8.4, RFC 2118 MPPC token grammar). A single +// instance is shared between the fast-path and slow-path receivers of one RDP +// connection. The token grammar follows the MS-RDPBCGR pseudo-code as +// implemented in FreeRDP's libfreerdp/codec/mppc.c (rdesktop's variant uses a +// different literal/length encoding that does not interop with Win10 RDP5 +// streams). +// +// Token grammar (bit-level, MSB first): +// +// literal 0x00-0x7F : "0" + 7 bits (8-bit token) +// literal 0x80-0xFF : "10" + 7 bits (9-bit token) +// copy tuple : "11" + offset prefix + length prefix +// +// Offset prefix (RDP5, compressionType nibble 0x1, 64K dictionary): +// +// "11111" + 6 bits → offset 0-63 +// "11110" + 8 bits → offset 64-319 +// "1110" + 11 bits → offset 320-2367 +// "110" + 16 bits → offset 2368-67903 +// +// (RDP4 / 8K dictionary: "1111"+6, "1110"+8, "110"+13.) +// +// Length prefix ("0" → 3; otherwise n 1-bits, a 0 terminator, then n+1 value +// bits with the (1<= 8 { + if br.readBit() == 0 { + // Literal 0x00-0x7F ("0" + 7 bits). + if d.offset >= mppcHistorySize { + return abort("history full") + } + d.history[d.offset] = byte(br.readBits(7)) + d.offset++ + continue + } + if br.readBit() == 0 { + // Literal 0x80-0xFF ("10" + 7 bits, 9-bit token). + if d.offset >= mppcHistorySize { + return abort("history full") + } + d.history[d.offset] = byte(0x80 | br.readBits(7)) + d.offset++ + continue + } + + // Copy tuple: decode CopyOffset (distance back from the write + // cursor, masked into the dictionary)。前缀位必须逐位惰性读取, + // 不能在 switch 初始化里预先消耗。 + var copyOffset int + if br.readBit() == 0 { + copyOffset = br.readBits(16) + 2368 + } else if br.readBit() == 0 { + copyOffset = br.readBits(11) + 320 + } else if br.readBit() == 0 { + copyOffset = br.readBits(8) + 64 + } else { + copyOffset = br.readBits(6) + } + + // Decode LengthOfMatch: n leading 1-bits + 0 terminator, then n+1 + // value bits; length = (1<<(n+1)) | bits. "0" alone means 3. + n := 0 + for br.readBit() == 1 { + n++ + const maxBits = 15 + if n > maxBits-1 { + return abort("length code overflow") + } + } + var copyLength int + if n == 0 { + copyLength = 3 + } else { + m := n + 1 + copyLength = (1 << uint(m)) | br.readBits(m) + } + + if d.offset+copyLength > mppcHistorySize { + return abort(fmt.Sprintf("copy overflows history (%d+%d)", d.offset, copyLength)) + } + src := (d.offset - copyOffset) & mask + for i := 0; i < copyLength; i++ { + d.history[d.offset] = d.history[src] + d.offset++ + src = (src + 1) & mask + } + } + + out := make([]byte, d.offset-start) + copy(out, d.history[start:d.offset]) + return out, nil +} + +// mppcBitReader reads bits MSB-first from a byte slice. Reading past the +// end yields zero bits (matching FreeRDP's padded bit stream behaviour). +type mppcBitReader struct { + data []byte + byteIdx int + mask byte // bit mask within current byte; starts at 0x80 + bitsLeft int // total bits remaining +} + +func newMppcBitReader(data []byte) *mppcBitReader { + return &mppcBitReader{ + data: data, + mask: 0x80, + bitsLeft: len(data) * 8, + } +} + +func (r *mppcBitReader) readBit() int { + if r.bitsLeft <= 0 { + return 0 + } + r.bitsLeft-- + var bit int + if r.data[r.byteIdx]&r.mask != 0 { + bit = 1 + } + r.mask >>= 1 + if r.mask == 0 { + r.mask = 0x80 + r.byteIdx++ + } + return bit +} + +func (r *mppcBitReader) readBits(n int) int { + result := 0 + for i := 0; i < n; i++ { + result = (result << 1) | r.readBit() + } + return result +} diff --git a/core/mppc_test.go b/core/mppc_test.go new file mode 100644 index 0000000..b41c9d6 --- /dev/null +++ b/core/mppc_test.go @@ -0,0 +1,353 @@ +package core + +import ( + "bytes" + "math/rand" + "testing" +) + +// Token grammar reference (MS-RDPBCGR §3.1.8.4 as implemented by FreeRDP's +// libfreerdp/codec/mppc.c): +// +// literal 0x00-0x7F : "0" + 7 bits (8-bit token) +// literal 0x80-0xFF : "10" + 7 bits (9-bit token) +// copy tuple : "11" + offset prefix + length prefix + +// compressedABC is the MPPC-64K encoding of "abc" (three literal tokens; +// bytes < 0x80 are transmitted as 8-bit "0"+7bits tokens). +var compressedABC = []byte{0x61, 0x62, 0x63} + +// compressedABCABC is "abc" followed by a copy tuple that repeats it: +// +// "11" copy tuple, "111" offset group 1, offset 000011 (=3), +// length "0" (=3). Bits: 01100001 01100010 01100011 11111000 0110… +// → 0x61 0x62 0x63 0xF8 0x60 +// +// Decoded: "abcabc". +var compressedABCABC = []byte{0x61, 0x62, 0x63, 0xF8, 0x60} + +func TestMppcDecompressLiterals(t *testing.T) { + d := NewMppcDecompressor() + got, err := d.Decompress(mppcType64K|mppcCompressed, compressedABC) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, []byte("abc")) { + t.Errorf("got %q, want %q", got, "abc") + } +} + +func TestMppcDecompressCopyTuple(t *testing.T) { + d := NewMppcDecompressor() + got, err := d.Decompress(mppcType64K|mppcCompressed, compressedABCABC) + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, []byte("abcabc")) { + t.Errorf("got %q, want %q", got, "abcabc") + } +} + +func TestMppcDecompressFlush(t *testing.T) { + d := NewMppcDecompressor() + // First call: populate history. + _, _ = d.Decompress(mppcType64K|mppcCompressed, compressedABC) + if d.history[0] != 'a' { + t.Fatal("history not populated after first call") + } + // Second call with FLUSH: history must be zeroed, offset restarts at 0. + _, _ = d.Decompress(mppcAtFront|mppcFlushed|mppcCompressed, compressedABC) + if d.history[0] != 'a' || d.history[1] != 'b' || d.history[2] != 'c' { + t.Errorf("unexpected history after flush: %q %q %q", + d.history[0], d.history[1], d.history[2]) + } + if d.offset != 3 { + t.Errorf("offset after flush: got %d, want 3", d.offset) + } +} + +func TestMppcDecompressReset(t *testing.T) { + d := NewMppcDecompressor() + _, _ = d.Decompress(mppcType64K|mppcCompressed, compressedABC) + if d.offset != 3 { + t.Fatalf("offset after first call: got %d, want 3", d.offset) + } + // RESET (PACKET_AT_FRONT): cursor returns to the front, history contents + // preserved; the fresh tokens overwrite from position 0. + _, _ = d.Decompress(mppcAtFront|mppcType64K|mppcCompressed, compressedABC) + if d.offset != 3 { + t.Fatalf("offset after reset+decompress: got %d, want 3", d.offset) + } + if d.history[0] != 'a' { + t.Errorf("history[0] after reset: got %q", d.history[0]) + } +} + +func TestMppcDecompressUncompressed(t *testing.T) { + d := NewMppcDecompressor() + plain := []byte("hello") + got, err := d.Decompress(mppcType64K, plain) // no COMPRESSED flag + if err != nil { + t.Fatal(err) + } + if !bytes.Equal(got, plain) { + t.Errorf("got %q, want %q", got, plain) + } + // Uncompressed segments do NOT advance the shared history (matching + // FreeRDP): only compressed tokens feed it. + if d.offset != 0 { + t.Errorf("offset: got %d, want 0 (history untouched)", d.offset) + } +} + +// mppcTestEncoder is a minimal MPPC compressor used to generate reference +// streams. It mirrors the decoder's grammar and history arithmetic exactly. +type mppcTestEncoder struct { + buf bitWriter + hist []byte + mask int + big bool +} + +type bitWriter struct { + out []byte + cur byte + nbit uint +} + +func (w *bitWriter) writeBit(b int) { + if b != 0 { + w.cur |= 0x80 >> w.nbit + } + w.nbit++ + if w.nbit == 8 { + w.out = append(w.out, w.cur) + w.cur = 0 + w.nbit = 0 + } +} + +func (w *bitWriter) writeBits(v, n int) { + for i := n - 1; i >= 0; i-- { + w.writeBit((v >> uint(i)) & 1) + } +} + +func (w *bitWriter) flush() []byte { + if w.nbit > 0 { + w.out = append(w.out, w.cur) + w.cur = 0 + w.nbit = 0 + } + return w.out +} + +func newMppcTestEncoder(big bool) *mppcTestEncoder { + mask := 8191 + if big { + mask = 65535 + } + return &mppcTestEncoder{big: big, mask: mask} +} + +func (e *mppcTestEncoder) emitLiteral(b byte) { + if b < 0x80 { + e.buf.writeBits(int(b), 8) + } else { + e.buf.writeBits(2, 2) // "10" + e.buf.writeBits(int(b)&0x7F, 7) + } + e.hist = append(e.hist, b) +} + +func (e *mppcTestEncoder) emitCopy(off, length int) { + e.buf.writeBits(3, 2) // "11" + if e.big { + switch { + case off <= 63: + e.buf.writeBits(7, 3) // "111" + e.buf.writeBits(off, 6) + case off <= 319: + e.buf.writeBits(6, 3) // "110" + e.buf.writeBits(off-64, 8) + case off <= 2367: + e.buf.writeBits(2, 2) // "10" + e.buf.writeBits(off-320, 11) + default: + e.buf.writeBit(0) + e.buf.writeBits(off-2368, 16) + } + } else { + switch { + case off <= 63: + e.buf.writeBits(3, 2) // "11" + e.buf.writeBits(off, 6) + case off <= 319: + e.buf.writeBits(2, 2) // "10" + e.buf.writeBits(off-64, 8) + default: + e.buf.writeBit(0) + e.buf.writeBits(off-320, 13) + } + } + // Length: "0" → 3; otherwise n leading 1-bits, a 0 terminator, then n+1 + // value bits with the (1< len(enc.hist) { + b := byte(rng.Intn(256)) + enc.emitLiteral(b) + want = append(want, b) + } + enc.emitCopy(off, length) + pos := len(want) + src := (pos - off) & enc.mask + for i := 0; i < length; i++ { + want = append(want, want[src]) + src = (src + 1) & enc.mask + } + } + got, err := NewMppcDecompressor().Decompress(mppcType64K|mppcCompressed, enc.buf.flush()) + if err != nil { + t.Fatalf("%v", err) + } + if !bytes.Equal(got, want) { + t.Fatalf("copy round-trip mismatch (%d vs %d bytes)", len(got), len(want)) + } +} + +func TestMppcHistoryAcrossSegments(t *testing.T) { + enc := newMppcTestEncoder(true) + d := NewMppcDecompressor() + flags := byte(mppcCompressed) // 无 RESET/AT_FRONT:历史跨段连续 + + // Segment 1: seed literals. + var want []byte + for i := 0; i < 400; i++ { + b := byte('A' + i%26) + enc.emitLiteral(b) + want = append(want, b) + } + got, err := d.Decompress(flags, enc.buf.flush()) + if err != nil || !bytes.Equal(got, want[:len(got)]) { + t.Fatalf("seg1: err=%v len=%d", err, len(got)) + } + enc.buf.out = nil + + // Segment 2: copy referencing segment 1 content. + enc.emitCopy(200, 50) + got2, err := d.Decompress(flags, enc.buf.flush()) + if err != nil { + t.Fatalf("seg2: %v", err) + } + want2 := want[len(want)-200:] + for i := 0; i < 50; i++ { + if got2[i] != want2[i] { + t.Fatalf("seg2 mismatch at %d: got %q want %q", i, got2[i], want2[i]) + } + } +} + +func TestMppcFlushAndReset(t *testing.T) { + enc := newMppcTestEncoder(true) + d := NewMppcDecompressor() + flags := byte(mppcCompressed) // 无 RESET/AT_FRONT:历史跨段连续 + + for i := 0; i < 100; i++ { + enc.emitLiteral(byte(i % 256)) + } + if _, err := d.Decompress(flags, enc.buf.flush()); err != nil { + t.Fatal(err) + } + if d.offset == 0 { + t.Fatal("history should have advanced") + } + + // RESET moves the write cursor to the front without clearing content. + enc2 := newMppcTestEncoder(true) + enc2.emitLiteral('x') + enc2.emitLiteral('y') + got, err := d.Decompress(mppcCompressed|mppcAtFront|mppcType64K, enc2.buf.flush()) + if err != nil || string(got) != "xy" { + t.Fatalf("reset: err=%v got=%q", err, got) + } + if d.offset != 2 { + t.Fatalf("reset should restart the cursor, got %d", d.offset) + } + + // FLUSH clears the whole dictionary. + got, err = d.Decompress(mppcCompressed|mppcAtFront|mppcFlushed, enc2.buf.flush()) + if err != nil || string(got) != "xy" { + t.Fatalf("flush: err=%v got=%q", err, got) + } +} + +func TestMppcRejectsGarbage(t *testing.T) { + d := NewMppcDecompressor() + // A long run of 1-bits overflows the LengthOfMatch prefix code. + data := make([]byte, 16) + for i := range data { + data[i] = 0xFF + } + if _, err := d.Decompress(mppcCompressed|mppcType64K, data); err == nil { + t.Fatal("expected error for length code overflow") + } +} diff --git a/core/rle.go b/core/rle.go new file mode 100644 index 0000000..2235529 --- /dev/null +++ b/core/rle.go @@ -0,0 +1,997 @@ +package core + +import ( + "fmt" + "log/slog" + "sync" + "unsafe" +) + +func CVAL(p *[]uint8) int { + a := int((*p)[0]) + *p = (*p)[1:] + return a +} + +func CVAL2(p *[]uint8, v *uint16) { + *v = *((*uint16)(unsafe.Pointer(&(*p)[0]))) + *p = (*p)[2:] +} + +func CVAL3(p *[]uint8, v *[3]uint8) { + (*v)[0] = (*p)[0] + (*v)[1] = (*p)[1] + (*v)[2] = (*p)[2] + *p = (*p)[3:] +} + +func REPEAT(f func(), count *int, x *int, width int) { + for *count > 0 && *x < width { + f() + *count-- + *x++ + } +} + +// rleFailLog 记录 RLE 解码失败(限次),此前失败被静默吞掉: +// 半解码的池化缓冲被直接画上画布形成噪块,且无任何日志可查。 +func rleFailLog(reason string, width, height, inputLen int, consumed int) { + slog.Warn("bitmap RLE decode failed", "reason", reason, + "w", width, "h", height, "inputLen", inputLen, "consumed", consumed) +} + +/* 1 byte bitmap decompress */ +func decompress1(output *[]uint8, width, height int, input []uint8, size int) bool { + var ( + prevline, line, count int + offset, code int + x int = width + opcode int + lastopcode int8 = -1 + insertmix, bicolour, isfillormix bool + mixmask, mask uint8 + colour1, colour2 uint8 + mix uint8 = 0xff + fom_mask uint8 + ) + out := *output + for len(input) != 0 { + fom_mask = 0 + code = CVAL(&input) + opcode = code >> 4 + /* Handle different opcode forms */ + switch opcode { + case 0xc, 0xd, 0xe: + opcode -= 6 + count = int(code & 0xf) + offset = 16 + break + case 0xf: + opcode = code & 0xf + if opcode < 9 { + count = int(CVAL(&input)) + count |= int(CVAL(&input) << 8) + } else { + count = 1 + if opcode < 0xb { + count = 8 + } + } + offset = 0 + break + default: + opcode >>= 1 + count = int(code & 0x1f) + offset = 32 + break + } + /* Handle strange cases for counts */ + if offset != 0 { + isfillormix = ((opcode == 2) || (opcode == 7)) + if count == 0 { + if isfillormix { + count = int(CVAL(&input)) + 1 + } else { + count = int(CVAL(&input) + offset) + } + } else if isfillormix { + count <<= 3 + } + } + /* Read preliminary data */ + switch opcode { + case 0: /* Fill */ + if (lastopcode == int8(opcode)) && !((x == width) && (prevline == 0)) { + insertmix = true + } + break + case 8: /* Bicolour */ + colour1 = uint8(CVAL(&input)) + colour2 = uint8(CVAL(&input)) + break + case 3: /* Colour */ + colour2 = uint8(CVAL(&input)) + break + case 6: /* SetMix/Mix */ + fallthrough + case 7: /* SetMix/FillOrMix */ + mix = uint8(CVAL(&input)) + opcode -= 5 + break + case 9: /* FillOrMix_1 */ + mask = 0x03 + opcode = 0x02 + fom_mask = 3 + break + case 0x0a: /* FillOrMix_2 */ + mask = 0x05 + opcode = 0x02 + fom_mask = 5 + break + } + lastopcode = int8(opcode) + mixmask = 0 + /* Output body */ + for count > 0 { + if x >= width { + if height <= 0 { + return false + } + + x = 0 + height-- + prevline = line + line = height * width + } + switch opcode { + case 0: /* Fill */ + if insertmix { + if prevline == 0 { + out[x+line] = mix + } else { + out[x+line] = out[prevline+x] ^ mix + } + insertmix = false + count-- + x++ + } + n := min(count, width-x) + if prevline == 0 { + clear(out[x+line : x+line+n]) + } else { + copy(out[x+line:x+line+n], out[prevline+x:prevline+x+n]) + } + count -= n + x += n + break + case 1: /* Mix */ + n := min(count, width-x) + if prevline == 0 { + seg := out[x+line : x+line+n] + for i := range seg { + seg[i] = mix + } + } else { + src := out[prevline+x : prevline+x+n] + dst := out[x+line : x+line+n] + for i := range dst { + dst[i] = src[i] ^ mix + } + } + count -= n + x += n + break + case 2: /* Fill or Mix */ + if prevline == 0 { + for count > 0 && x < width { + mixmask <<= 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[x+line] = mix + } else { + out[x+line] = 0 + } + + count-- + x++ + } + } else { + for count > 0 && x < width { + mixmask = mixmask << 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[x+line] = out[prevline+x] ^ mix + } else { + out[x+line] = out[prevline+x] + } + + count-- + x++ + } + } + break + case 3: /* Colour */ + n := min(count, width-x) + seg := out[x+line : x+line+n] + for i := range seg { + seg[i] = colour2 + } + count -= n + x += n + break + case 4: /* Copy */ + for count > 0 && x < width { + n := min(count, width-x) + if len(input) < n { + return false + } + copy(out[x+line:x+line+n], input[:n]) + input = input[n:] + count -= n + x += n + } + break + case 8: /* Bicolour */ + for count > 0 && x < width { + if bicolour { + out[x+line] = colour2 + bicolour = false + } else { + out[x+line] = colour1 + bicolour = true + count++ + } + + count-- + x++ + } + + break + + case 0xd: /* White */ + n := min(count, width-x) + seg := out[x+line : x+line+n] + for i := range seg { + seg[i] = 0xff + } + count -= n + x += n + break + case 0xe: /* Black */ + n := min(count, width-x) + clear(out[x+line : x+line+n]) + count -= n + x += n + break + default: + fmt.Printf("bitmap opcode 0x%x\n", opcode) + return false + } + } + } + return true +} + +// decompress2Pool reuses the intermediate []uint16 buffer across calls to +// decompress2, avoiding a large per-frame allocation. +var decompress2Pool sync.Pool + +/* 2 byte bitmap decompress */ +func decompress2(output *[]uint8, width, height int, input []uint8, size int) bool { + needed := width * height + + var out []uint16 + if v := decompress2Pool.Get(); v != nil { + out = v.([]uint16) + if cap(out) < needed { + out = make([]uint16, needed) + } else { + out = out[:needed] + } + } else { + out = make([]uint16, needed) + } + defer func() { decompress2Pool.Put(out[:cap(out)]) }() + + var ( + prevline, line, count int + offset, code int + x int = width + opcode int + lastopcode int = -1 + insertmix, bicolour, isfillormix bool + mixmask, mask uint8 + colour1, colour2 uint16 + mix uint16 = 0xffff + fom_mask uint8 + ) + inputLen0 := len(input) + + for len(input) != 0 { + fom_mask = 0 + code = CVAL(&input) + opcode = code >> 4 + /* Handle different opcode forms */ + switch opcode { + case 0xc, 0xd, 0xe: + opcode -= 6 + count = code & 0xf + offset = 16 + break + case 0xf: + opcode = code & 0xf + if opcode < 9 { + count = CVAL(&input) + count |= CVAL(&input) << 8 + } else { + count = 1 + if opcode < 0xb { + count = 8 + } + } + offset = 0 + break + default: + opcode >>= 1 + count = code & 0x1f + offset = 32 + break + } + + /* Handle strange cases for counts */ + if offset != 0 { + isfillormix = ((opcode == 2) || (opcode == 7)) + if count == 0 { + if isfillormix { + count = CVAL(&input) + 1 + } else { + count = CVAL(&input) + offset + } + } else if isfillormix { + count <<= 3 + } + } + /* Read preliminary data */ + switch opcode { + case 0: /* Fill */ + if (lastopcode == opcode) && !((x == width) && (prevline == 0)) { + insertmix = true + } + break + case 8: /* Bicolour */ + CVAL2(&input, &colour1) + CVAL2(&input, &colour2) + break + case 3: /* Colour */ + CVAL2(&input, &colour2) + break + case 6: /* SetMix/Mix */ + fallthrough + case 7: /* SetMix/FillOrMix */ + CVAL2(&input, &mix) + opcode -= 5 + break + case 9: /* FillOrMix_1 */ + mask = 0x03 + opcode = 0x02 + fom_mask = 3 + break + case 0x0a: /* FillOrMix_2 */ + mask = 0x05 + opcode = 0x02 + fom_mask = 5 + break + } + lastopcode = opcode + mixmask = 0 + /* Output body */ + for count > 0 { + if x >= width { + if height <= 0 { + rleFailLog("lines exhausted", width, height, inputLen0, inputLen0-len(input)) + return false + } + + x = 0 + height-- + prevline = line + line = height * width + } + switch opcode { + case 0: /* Fill */ + if insertmix { + if prevline == 0 { + out[x+line] = mix + } else { + out[x+line] = out[prevline+x] ^ mix + } + insertmix = false + count-- + x++ + } + n := min(count, width-x) + if prevline == 0 { + clear(out[x+line : x+line+n]) + } else { + copy(out[x+line:x+line+n], out[prevline+x:prevline+x+n]) + } + count -= n + x += n + break + case 1: /* Mix */ + n := min(count, width-x) + if prevline == 0 { + seg := out[x+line : x+line+n] + for i := range seg { + seg[i] = mix + } + } else { + src := out[prevline+x : prevline+x+n] + dst := out[x+line : x+line+n] + for i := range dst { + dst[i] = src[i] ^ mix + } + } + count -= n + x += n + break + case 2: /* Fill or Mix */ + if prevline == 0 { + for count > 0 && x < width { + mixmask <<= 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[x+line] = mix + } else { + out[x+line] = 0 + } + + count-- + x++ + } + } else { + for count > 0 && x < width { + mixmask = mixmask << 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[x+line] = out[prevline+x] ^ mix + } else { + out[x+line] = out[prevline+x] + } + + count-- + x++ + } + } + break + case 3: /* Colour */ + n := min(count, width-x) + seg := out[x+line : x+line+n] + for i := range seg { + seg[i] = colour2 + } + count -= n + x += n + break + case 4: /* Copy */ + for count > 0 && x < width { + n := min(count, width-x) + if len(input) < n*2 { + rleFailLog("copy input exhausted", width, height, inputLen0, inputLen0-len(input)) + return false + } + copy(out[x+line:x+line+n], unsafe.Slice((*uint16)(unsafe.Pointer(&input[0])), n)) + input = input[n*2:] + count -= n + x += n + } + + break + case 8: /* Bicolour */ + for count > 0 && x < width { + if bicolour { + out[x+line] = colour2 + bicolour = false + } else { + out[x+line] = colour1 + bicolour = true + count++ + } + + count-- + x++ + } + + break + case 0xd: /* White */ + n2 := min(count, width-x) + seg2 := out[x+line : x+line+n2] + for i := range seg2 { + seg2[i] = 0xffff + } + count -= n2 + x += n2 + break + case 0xe: /* Black */ + n3 := min(count, width-x) + clear(out[x+line : x+line+n3]) + count -= n3 + x += n3 + break + default: + rleFailLog(fmt.Sprintf("bad opcode 0x%x", opcode), width, height, inputLen0, inputLen0-len(input)) + return false + } + } + } + outBytes := *output + for i, v := range out { + outBytes[i*2] = byte(v >> 8) + outBytes[i*2+1] = byte(v) + } + return true +} + +// /* 3 byte bitmap decompress */ +func decompress3(output *[]uint8, width, height int, input []uint8, size int) bool { + var ( + prevline, line, count int + opcode, offset, code int + x int = width + lastopcode int = -1 + insertmix, bicolour, isfillormix bool + mixmask, mask uint8 + colour1 = [3]uint8{0, 0, 0} + colour2 = [3]uint8{0, 0, 0} + mix = [3]uint8{0xff, 0xff, 0xff} + fom_mask uint8 + ) + out := *output + for len(input) != 0 { + fom_mask = 0 + code = CVAL(&input) + opcode = code >> 4 + /* Handle different opcode forms */ + switch opcode { + case 0xc, 0xd, 0xe: + opcode -= 6 + count = code & 0xf + offset = 16 + break + case 0xf: + opcode = code & 0xf + if opcode < 9 { + count = CVAL(&input) + count |= CVAL(&input) << 8 + } else { + count = 1 + if opcode < 0xb { + count = 8 + } + } + offset = 0 + break + default: + opcode >>= 1 + count = code & 0x1f + offset = 32 + break + } + + /* Handle strange cases for counts */ + if offset != 0 { + isfillormix = ((opcode == 2) || (opcode == 7)) + if count == 0 { + if isfillormix { + count = CVAL(&input) + 1 + } else { + count = CVAL(&input) + offset + } + } else if isfillormix { + count <<= 3 + } + } + /* Read preliminary data */ + switch opcode { + case 0: /* Fill */ + if (lastopcode == opcode) && !((x == width) && (prevline == 0)) { + insertmix = true + } + break + case 8: /* Bicolour */ + CVAL3(&input, &colour1) + CVAL3(&input, &colour2) + break + case 3: /* Colour */ + CVAL3(&input, &colour2) + break + case 6: /* SetMix/Mix */ + fallthrough + case 7: /* SetMix/FillOrMix */ + CVAL3(&input, &mix) + opcode -= 5 + break + case 9: /* FillOrMix_1 */ + mask = 0x03 + opcode = 0x02 + fom_mask = 3 + break + case 0x0a: /* FillOrMix_2 */ + mask = 0x05 + opcode = 0x02 + fom_mask = 5 + break + } + + lastopcode = opcode + mixmask = 0 + /* Output body */ + for count > 0 { + if x >= width { + if height <= 0 { + return false + } + + x = 0 + height-- + prevline = line + line = height * width * 3 + } + switch opcode { + case 0: /* Fill */ + if insertmix { + if prevline == 0 { + out[3*x+line] = mix[0] + out[3*x+line+1] = mix[1] + out[3*x+line+2] = mix[2] + } else { + out[3*x+line] = out[prevline+3*x] ^ mix[0] + out[3*x+line+1] = out[prevline+3*x+1] ^ mix[1] + out[3*x+line+2] = out[prevline+3*x+2] ^ mix[2] + } + insertmix = false + count-- + x++ + } + n := min(count, width-x) + if prevline == 0 { + clear(out[3*x+line : 3*x+line+3*n]) + } else { + dstBase := 3*x + line + srcBase := prevline + 3*x + copy(out[dstBase:dstBase+3*n], out[srcBase:srcBase+3*n]) + } + count -= n + x += n + break + case 1: /* Mix */ + n := min(count, width-x) + dst1 := out[3*x+line : 3*x+line+3*n] + if prevline == 0 { + // Exponential-doubling copy: O(log n) memcpy calls + dst1[0], dst1[1], dst1[2] = mix[0], mix[1], mix[2] + for wrote := 3; wrote < len(dst1); { + wrote += copy(dst1[wrote:], dst1[:wrote]) + } + } else { + src1 := out[prevline+3*x : prevline+3*x+3*n] + for i := 0; i+2 < len(dst1); i += 3 { + dst1[i] = src1[i] ^ mix[0] + dst1[i+1] = src1[i+1] ^ mix[1] + dst1[i+2] = src1[i+2] ^ mix[2] + } + } + count -= n + x += n + break + case 2: /* Fill or Mix */ + if prevline == 0 { + base := 3*x + line + for count > 0 && x < width { + mixmask = mixmask << 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[base] = mix[0] + out[base+1] = mix[1] + out[base+2] = mix[2] + } else { + out[base] = 0 + out[base+1] = 0 + out[base+2] = 0 + } + base += 3 + count-- + x++ + } + } else { + base := 3*x + line + prev := prevline + 3*x + for count > 0 && x < width { + mixmask = mixmask << 1 + if mixmask == 0 { + mask = fom_mask + if fom_mask == 0 { + mask = uint8(CVAL(&input)) + mixmask = 1 + } + } + if mask&mixmask != 0 { + out[base] = out[prev] ^ mix[0] + out[base+1] = out[prev+1] ^ mix[1] + out[base+2] = out[prev+2] ^ mix[2] + } else { + out[base] = out[prev] + out[base+1] = out[prev+1] + out[base+2] = out[prev+2] + } + base += 3 + prev += 3 + count-- + x++ + } + } + break + case 3: /* Colour */ + n := min(count, width-x) + seg3 := out[3*x+line : 3*x+line+3*n] + seg3[0], seg3[1], seg3[2] = colour2[0], colour2[1], colour2[2] + for wrote := 3; wrote < len(seg3); { + wrote += copy(seg3[wrote:], seg3[:wrote]) + } + count -= n + x += n + break + case 4: /* Copy */ + for count > 0 && x < width { + n := min(count, width-x) + if len(input) < n*3 { + return false + } + copy(out[3*x+line:3*x+line+n*3], input[:n*3]) + input = input[n*3:] + count -= n + x += n + } + break + case 8: /* Bicolour */ + base8 := 3*x + line + for count > 0 && x < width { + if bicolour { + out[base8] = colour2[0] + out[base8+1] = colour2[1] + out[base8+2] = colour2[2] + bicolour = false + } else { + out[base8] = colour1[0] + out[base8+1] = colour1[1] + out[base8+2] = colour1[2] + bicolour = true + count++ + } + base8 += 3 + count-- + x++ + } + break + case 0xd: /* White */ + n3 := min(count, width-x) + seg3 := out[3*x+line : 3*x+line+3*n3] + for i := range seg3 { + seg3[i] = 0xff + } + count -= n3 + x += n3 + break + case 0xe: /* Black */ + n2 := min(count, width-x) + clear(out[3*x+line : 3*x+line+3*n2]) + count -= n2 + x += n2 + break + default: + fmt.Printf("bitmap opcode 0x%x\n", opcode) + return false + } + } + } + + return true +} + +/* decompress a colour plane */ +func processPlane(in *[]uint8, width, height int, output *[]uint8, j int) int { + var ( + indexw int + indexh int + code int + collen int + replen int + color uint8 + x uint8 + revcode int + lastline int + thisline int + ) + ln := len(*in) + out := *output // hoist pointer dereference; writes to out[i] affect the underlying array + + lastline = 0 + indexh = 0 + i := 0 + for indexh < height { + thisline = j + (width * height * 4) - ((indexh + 1) * width * 4) + color = 0 + indexw = 0 + i = thisline + + if lastline == 0 { + for indexw < width { + code = CVAL(in) + replen = int(code & 0xf) + collen = int((code >> 4) & 0xf) + revcode = (replen << 4) | collen + if (revcode <= 47) && (revcode >= 16) { + replen = revcode + collen = 0 + } + for collen > 0 { + color = uint8(CVAL(in)) + out[i] = color + i += 4 + + indexw++ + collen-- + } + for replen > 0 { + out[i] = color + i += 4 + indexw++ + replen-- + } + } + } else { + // prevOffset is constant per row: indexw*4+lastline == i+(lastline-thisline) + // because i == thisline+indexw*4. Pre-computing it eliminates a multiply + // per pixel in both inner loops. + prevOffset := lastline - thisline + for indexw < width { + code = CVAL(in) + replen = int(code & 0xf) + collen = int((code >> 4) & 0xf) + revcode = (replen << 4) | collen + if (revcode <= 47) && (revcode >= 16) { + replen = revcode + collen = 0 + } + for collen > 0 { + x = uint8(CVAL(in)) + if x&1 != 0 { + x = x >> 1 + x = x + 1 + color = -x + } else { + x = x >> 1 + color = x + } + x = out[i+prevOffset] + color + out[i] = x + i += 4 + indexw++ + collen-- + } + for replen > 0 { + x = out[i+prevOffset] + color + out[i] = x + i += 4 + indexw++ + replen-- + } + } + } + indexh++ + lastline = thisline + } + return ln - len(*in) +} + +/* 4 byte bitmap decompress */ +func decompress4(output *[]uint8, width, height int, input []uint8, size int) bool { + var ( + code int + onceBytes, total int + ) + + code = CVAL(&input) + rle := code&0x10 != 0 + noAlpha := code&0x20 != 0 + + if !rle { + return false + } + + total = 1 + out := *output + + if noAlpha { + // No alpha plane in the stream; fill alpha channel with 0xFF. + for i := 3; i < len(out); i += 4 { + out[i] = 0xFF + } + } else { + onceBytes = processPlane(&input, width, height, output, 3) + total += onceBytes + } + + onceBytes = processPlane(&input, width, height, output, 2) + total += onceBytes + + onceBytes = processPlane(&input, width, height, output, 1) + total += onceBytes + + onceBytes = processPlane(&input, width, height, output, 0) + total += onceBytes + + return true +} + +// DecompressInto decompresses bitmap data into dst, reusing dst if it has +// sufficient capacity (size = width*height*bpp). If dst is nil or too small +// a new slice is allocated. Returns the (re)used output slice. +func DecompressInto(input []uint8, dst []uint8, width, height int, bpp int) ([]uint8, bool) { + size := width * height * bpp + if cap(dst) >= size { + dst = dst[:size] + } else { + dst = make([]uint8, size) + } + ok := false + switch bpp { + case 1: + ok = decompress1(&dst, width, height, input, size) + case 2: + ok = decompress2(&dst, width, height, input, size) + case 3: + ok = decompress3(&dst, width, height, input, size) + case 4: + ok = decompress4(&dst, width, height, input, size) + default: + fmt.Printf("bpp %d\n", bpp) + } + return dst, ok +} + +/* main decompress function */ +func Decompress(input []uint8, width, height int, bpp int) []uint8 { + out, _ := DecompressInto(input, nil, width, height, bpp) + return out +} diff --git a/core/rle_test.go b/core/rle_test.go new file mode 100644 index 0000000..bc74ca2 --- /dev/null +++ b/core/rle_test.go @@ -0,0 +1,26 @@ +// rle_test.go +package core + +import ( + "fmt" + "testing" +) + +func BenchmarkDecompress3(b *testing.B) { + input := []byte{ + 192, 44, 200, 8, 132, 200, 8, 200, 8, 200, 8, 200, 8, 0, 19, 132, 232, 8, 12, 50, 142, 66, 77, 58, 208, 59, 225, 25, 1, 0, 0, 0, 0, 0, 0, 0, 132, 139, 33, 142, 66, 142, 66, 142, 66, 208, 59, 4, 43, 1, 0, 0, 0, 0, 0, 0, 0, 132, 203, 41, 142, 66, 142, 66, 142, 66, 208, 59, 96, 0, 1, 0, 0, 0, 0, 0, 0, 0, 132, 9, 17, 142, 66, 142, 66, 142, 66, 208, 59, 230, 27, 1, 0, 0, 0, 0, 0, 0, 0, 132, 200, 8, 9, 17, 139, 33, 74, 25, 243, 133, 14, 200, 8, 132, 200, 8, 200, 8, 200, 8, 200, 8, + } + dst := make([]uint8, 64*64*3) + b.ResetTimer() + for i := 0; i < b.N; i++ { + DecompressInto(input, dst, 64, 64, 3) + } +} + +func TestSum(t *testing.T) { + input := []byte{ + 192, 44, 200, 8, 132, 200, 8, 200, 8, 200, 8, 200, 8, 0, 19, 132, 232, 8, 12, 50, 142, 66, 77, 58, 208, 59, 225, 25, 1, 0, 0, 0, 0, 0, 0, 0, 132, 139, 33, 142, 66, 142, 66, 142, 66, 208, 59, 4, 43, 1, 0, 0, 0, 0, 0, 0, 0, 132, 203, 41, 142, 66, 142, 66, 142, 66, 208, 59, 96, 0, 1, 0, 0, 0, 0, 0, 0, 0, 132, 9, 17, 142, 66, 142, 66, 142, 66, 208, 59, 230, 27, 1, 0, 0, 0, 0, 0, 0, 0, 132, 200, 8, 9, 17, 139, 33, 74, 25, 243, 133, 14, 200, 8, 132, 200, 8, 200, 8, 200, 8, 200, 8, + } + out := Decompress(input, 64, 64, 3) + fmt.Println(out) +} diff --git a/core/socket.go b/core/socket.go new file mode 100644 index 0000000..508a9fc --- /dev/null +++ b/core/socket.go @@ -0,0 +1,127 @@ +package core + +import ( + "bufio" + "crypto/rsa" + "crypto/sha256" + "crypto/tls" + "encoding/asn1" + "errors" + "math/big" + "net" + "time" +) + +// readBufSize is the size of the buffered reader used for socket reads. +// RDP packets can be large (bitmap updates, channel data); a 64 KiB buffer +// keeps the number of read(2) syscalls low without wasting memory. +const readBufSize = 65536 + +// tcpRecvBufSize is the OS-level TCP receive socket buffer size. +// The default on most systems (~87 KiB on Linux, ~128 KiB on macOS) is too +// small for high-resolution RDP sessions where the server can burst several +// MiB of bitmap/H.264 data per frame. 512 KiB allows the kernel to buffer +// more in-flight data, reducing stalls when the application goroutine is +// briefly busy decoding a previous frame. +const tcpRecvBufSize = 512 * 1024 + +type SocketLayer struct { + conn net.Conn + tlsConn *tls.Conn + reader *bufio.Reader // buffers reads regardless of TLS state + serverName string + // certVerifier 非空时在 TLS 握手完成后以服务器叶子证书的 SHA-256 指纹调用, + // 返回错误则中断连接(TOFU 指纹校验用) + certVerifier func(sha256Fp []byte) error +} + +func NewSocketLayer(conn net.Conn, serverName string) *SocketLayer { + // Disable Nagle's algorithm so small DVC responses are sent immediately. + if tc, ok := conn.(*net.TCPConn); ok { + tc.SetNoDelay(true) + // Increase the OS receive buffer so the kernel can absorb large bitmap + // or H.264 bursts without dropping bytes while the decoder is busy. + // SetReadBuffer is a best-effort hint; ignore errors (e.g. restricted + // by the OS cap in /proc/sys/net/core/rmem_max on Linux). + _ = tc.SetReadBuffer(tcpRecvBufSize) + } + l := &SocketLayer{ + conn: conn, + tlsConn: nil, + serverName: serverName, + } + l.reader = bufio.NewReaderSize(conn, readBufSize) + return l +} + +func (s *SocketLayer) SetDeadline(t time.Time) error { + return s.conn.SetDeadline(t) +} + +func (s *SocketLayer) Read(b []byte) (n int, err error) { + return s.reader.Read(b) +} + +func (s *SocketLayer) Write(b []byte) (n int, err error) { + if s.tlsConn != nil { + return s.tlsConn.Write(b) + } + return s.conn.Write(b) +} + +func (s *SocketLayer) Close() error { + if s.tlsConn != nil { + s.tlsConn.Close() // best-effort; always close the underlying TCP socket + } + return s.conn.Close() +} + +// SetCertVerifier 注册服务器证书指纹校验回调(须在 StartTLS 前调用) +func (s *SocketLayer) SetCertVerifier(fn func(sha256Fp []byte) error) { + s.certVerifier = fn +} + +func (s *SocketLayer) StartTLS() error { + config := &tls.Config{ + InsecureSkipVerify: true, + ServerName: s.serverName, + MinVersion: tls.VersionTLS12, + MaxVersion: tls.VersionTLS12, + // MaxVersion: tls.VersionTLS13, + } + tlsConn := tls.Client(s.conn, config) + if err := tlsConn.Handshake(); err != nil { + return err + } + // RDP 服务器普遍使用自签证书,链校验关闭;改为 TOFU 指纹校验: + // 应用层比对叶子证书 SHA-256,不匹配(如中间人)则拒绝继续 + if s.certVerifier != nil { + certs := tlsConn.ConnectionState().PeerCertificates + if len(certs) > 0 { + summary := sha256.Sum256(certs[0].Raw) + if err := s.certVerifier(summary[:]); err != nil { + s.tlsConn = tlsConn + return err + } + } + } + s.tlsConn = tlsConn + // Reset the buffered reader to read from the TLS connection. + // Reset discards any unconsumed buffered bytes from the plain-text phase, + // which is correct because the TLS handshake has already consumed them. + s.reader.Reset(tlsConn) + return nil +} + +type PublicKey struct { + N *big.Int `asn1:"explicit,tag:0"` // modulus + E int `asn1:"explicit,tag:1"` // public exponent +} + +func (s *SocketLayer) TlsPubKey() ([]byte, error) { + if s.tlsConn == nil { + return nil, errors.New("TLS conn does not exist") + } + pub := s.tlsConn.ConnectionState().PeerCertificates[0].PublicKey.(*rsa.PublicKey) + return asn1.Marshal(*pub) +} diff --git a/core/types.go b/core/types.go new file mode 100644 index 0000000..86120d3 --- /dev/null +++ b/core/types.go @@ -0,0 +1,26 @@ +package core + +import "git.zeroonesoft.cn/golib/rdplib/emission" + +type Transport interface { + Read(b []byte) (n int, err error) + Write(b []byte) (n int, err error) + Close() error + + On(event, listener any) *emission.Emitter + Once(event, listener any) *emission.Emitter + Off(event, listener any) *emission.Emitter + Emit(event any, arguments ...any) *emission.Emitter +} + +type FastPathListener interface { + RecvFastPath(secFlag byte, s []byte) +} + +type FastPathSender interface { + SendFastPath(secFlag byte, s []byte) (int, error) +} + +type ChannelSender interface { + SendToChannel(channel string, s []byte) (int, error) +} diff --git a/core/util.go b/core/util.go new file mode 100644 index 0000000..340ad04 --- /dev/null +++ b/core/util.go @@ -0,0 +1,53 @@ +package core + +import ( + "crypto/rand" + "encoding/binary" + "unicode/utf16" +) + +func Reverse(s []byte) []byte { + for i, j := 0, len(s)-1; i < j; i, j = i+1, j-1 { + s[i], s[j] = s[j], s[i] + } + return s +} + +func Random(n int) []byte { + const alpha = "0123456789ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz" + var bytes = make([]byte, n) + rand.Read(bytes) + for i, b := range bytes { + bytes[i] = alpha[b%byte(len(alpha))] + } + return bytes +} + +func UTF16ToLittleEndianBytes(u []uint16) []byte { + b := make([]byte, 2*len(u)) + for index, value := range u { + binary.LittleEndian.PutUint16(b[index*2:], value) + } + return b +} + +func LittleEndianBytesToUTF16(u []byte) []uint16 { + b := make([]uint16, len(u)/2) + for i := range b { + b[i] = binary.LittleEndian.Uint16(u[i*2:]) + } + return b +} + +// s.encode('utf-16le') +func UnicodeEncode(p string) []byte { + return UTF16ToLittleEndianBytes(utf16.Encode([]rune(p))) +} + +func UnicodeDecode(p []byte) string { + return string(utf16.Decode(LittleEndianBytesToUTF16(p))) +} + +func BytesToUint64(b []byte) uint64 { + return binary.LittleEndian.Uint64(b) +} diff --git a/doc.go b/doc.go new file mode 100644 index 0000000..9c32f00 --- /dev/null +++ b/doc.go @@ -0,0 +1,36 @@ +// Package grdp 实现微软 RDP(远程桌面协议)的纯 Go 客户端协议栈。 +// +// 本包是 nakagami/grdp 的深度 fork(MIT,保留上游版权),由 webrdp +// 项目扩展维护:NLA/CredSSP 认证、RDPGFX/H.264 多形态回调、动态分辨率、 +// Unicode 输入、多形态剪贴板、音频、驱动器重定向(rdpdr)等能力 +// 均在原生 Go API 之下,同一套协议栈可同时编译为原生程序与 +// GOOS=js GOARCH=wasm 的浏览器模块。 +// +// # 基本用法 +// +// client := grdp.NewRdpClient("host:3389", 1280, 800, +// func(addr string) (net.Conn, error) { return net.Dial("tcp", addr) }) +// client.OnError(func(e error) { ... }). +// OnReady(func() { ... }). +// OnBitmap(func(bits []grdp.Bitmap) { ... }) +// err := client.Login(domain, user, password) +// +// 所有 On* 回调在协议栈读循环 goroutine 内同步触发,回调不得长时间 +// 阻塞;Login 阻塞直到能力协商完成(会话就绪以 OnReady 为准)。 +// +// # 渲染模型 +// +// 传统管线经 OnBitmap 交付 Bitmap(RGBA()/FillRGBA 解码);GFX 管线经 +// OnH264Raw/OnH264I420/OnH264NV12 交付 H.264 NAL 或原始平面——解码 +// 由宿主负责(浏览器 WebCodecs、FFmpeg 均可),纯 Go 构建即可使用。 +// +// # 虚拟通道 +// +// 内置 rdpsnd/cliprdr/rdpdr/rdpgfx/rdpedisp/drdynvc(plugin 子包), +// 自定义通道实现 plugin.ChannelTransport 三方法后经 Channels.Register +// 注册,见 plugin/channel.go。 +// +// 目录导览:core/(传输抽象与工具)、emission/(事件发射器)、 +// protocol/(x224/tpkt/t125/sec/nla/pdu/lic/gcc 协议层)、 +// plugin/(虚拟通道框架与内置通道)。 +package grdp diff --git a/emission/emitter.go b/emission/emitter.go new file mode 100644 index 0000000..c2e3789 --- /dev/null +++ b/emission/emitter.go @@ -0,0 +1,304 @@ +// Package emission provides an event emitter. +// copy form https://raw.githubusercontent.com/chuckpreslar/emission/master/emitter.go +// fix issue with nest once +// +// Performance: common listener signatures (func(), func(error), func([]byte), +// func(uint16)) are dispatched without any reflection at emit time. +// Unknown signatures use a pre-built wrapper with pooled []reflect.Value to +// eliminate per-emit heap allocations. Generic helpers On1/Once1 allow +// callers to register any func(T) with zero reflection at emit time. + +package emission + +import ( + "errors" + "fmt" + "os" + "reflect" + "slices" + "sync" +) + +// Default number of maximum listeners for an event. +const DefaultMaxListeners = 10 + +// Error presented when an invalid argument is provided as a listener function +var ErrNoneFunction = errors.New("Kind of Value for listener is not Func.") + +// RecoveryListener ... +type RecoveryListener func(any, any, error) + +// listenerEntry pairs a pre-built dispatch wrapper with the original function +// pointer so that RemoveListener can identify and remove it. +type listenerEntry struct { + call func([]any) // reflection-free dispatch wrapper + ptr uintptr // reflect.ValueOf(original).Pointer(); 0 if unavailable +} + +// Pooled argument slices for the reflection fallback path. +// Indexed by argument count (0–3); counts > 3 allocate directly. +var rvPool = [4]*sync.Pool{ + 0: {New: func() any { v := make([]reflect.Value, 0); return &v }}, + 1: {New: func() any { v := make([]reflect.Value, 1); return &v }}, + 2: {New: func() any { v := make([]reflect.Value, 2); return &v }}, + 3: {New: func() any { v := make([]reflect.Value, 3); return &v }}, +} + +// Emitter is a reflect-minimal event emitter. +type Emitter struct { + events map[any][]listenerEntry + onces map[any][]listenerEntry + recoverer RecoveryListener + maxListeners int +} + +// NewEmitter returns a new Emitter object, defaulting the +// number of maximum listeners per event to the DefaultMaxListeners +// constant and initializing its events map. +func NewEmitter() *Emitter { + return &Emitter{ + events: make(map[any][]listenerEntry), + onces: make(map[any][]listenerEntry), + maxListeners: DefaultMaxListeners, + } +} + +// On is an alias for AddListener. +func (e *Emitter) On(event, listener any) *Emitter { + return e.AddListener(event, listener) +} + +// AddListener appends the listener argument to the event arguments slice. +// If the reflect Value of the listener does not have a Kind of Func then +// AddListener panics (or calls the RecoveryListener if one has been set). +func (e *Emitter) AddListener(event, listener any) *Emitter { + entry, ok := buildEntry(listener) + if !ok { + if e.recoverer == nil { + panic(ErrNoneFunction) + } + e.recoverer(event, listener, ErrNoneFunction) + return e + } + if e.maxListeners != -1 && e.maxListeners < len(e.events[event])+1 { + fmt.Fprintf(os.Stdout, "Warning: event `%v` has exceeded the maximum "+ + "number of listeners of %d.\n", event, e.maxListeners) + } + e.events[event] = append(e.events[event], entry) + return e +} + +// RemoveListener removes the listener from the event's listener slice. +func (e *Emitter) RemoveListener(event, listener any) *Emitter { + rv := reflect.ValueOf(listener) + if rv.Kind() != reflect.Func { + if e.recoverer == nil { + panic(ErrNoneFunction) + } + e.recoverer(event, listener, ErrNoneFunction) + return e + } + ptr := rv.Pointer() + if _, ok := e.events[event]; ok { + e.events[event] = slices.DeleteFunc(e.events[event], func(ent listenerEntry) bool { + return ent.ptr == ptr + }) + } + if _, ok := e.onces[event]; ok { + e.onces[event] = slices.DeleteFunc(e.onces[event], func(ent listenerEntry) bool { + return ent.ptr == ptr + }) + } + return e +} + +// Off is an alias for RemoveListener. +func (e *Emitter) Off(event, listener any) *Emitter { + return e.RemoveListener(event, listener) +} + +// Once registers a listener that fires at most once for the given event. +func (e *Emitter) Once(event, listener any) *Emitter { + entry, ok := buildEntry(listener) + if !ok { + if e.recoverer == nil { + panic(ErrNoneFunction) + } + e.recoverer(event, listener, ErrNoneFunction) + return e + } + if e.maxListeners != -1 && e.maxListeners < len(e.onces[event])+1 { + fmt.Fprintf(os.Stdout, "Warning: event `%v` has exceeded the maximum "+ + "number of listeners of %d.\n", event, e.maxListeners) + } + e.onces[event] = append(e.onces[event], entry) + return e +} + +// Emit calls each listener registered for event with the supplied arguments. +func (e *Emitter) Emit(event any, arguments ...any) *Emitter { + if entries, ok := e.events[event]; ok { + for _, ent := range entries { + e.dispatch(ent, event, arguments) + } + } + // Execute onces; preserve any new onces registered during execution + // (fix issue with nested Once — same semantics as original). + if entries, ok := e.onces[event]; ok { + origLen := len(entries) + for _, ent := range entries { + e.dispatch(ent, event, arguments) + } + e.onces[event] = e.onces[event][origLen:] + } + return e +} + +func (e *Emitter) dispatch(ent listenerEntry, event any, args []any) { + if e.recoverer != nil { + defer func() { + if r := recover(); r != nil { + e.recoverer(event, ent.ptr, fmt.Errorf("%v", r)) + } + }() + } + ent.call(args) +} + +// RecoverWith sets the listener to call when a panic occurs. +func (e *Emitter) RecoverWith(listener RecoveryListener) *Emitter { + e.recoverer = listener + return e +} + +// SetMaxListeners sets the maximum number of listeners per event. +// Pass -1 for unlimited. +func (e *Emitter) SetMaxListeners(max int) *Emitter { + e.maxListeners = max + return e +} + +// GetListenerCount returns the number of listeners registered for event. +func (e *Emitter) GetListenerCount(event any) (count int) { + if entries, ok := e.events[event]; ok { + count = len(entries) + } + return +} + +// On1 registers a typed listener with zero reflection at emit time. +// Use instead of e.On(event, fn) when the argument type is not covered by the +// built-in fast paths (func(), func(error), func([]byte), func(uint16)). +func On1[T any](e *Emitter, event any, fn func(T)) *Emitter { + e.events[event] = append(e.events[event], listenerEntry{ + call: func(args []any) { + if len(args) > 0 { + fn(args[0].(T)) + } + }, + }) + return e +} + +// Once1 registers a typed one-shot listener with zero reflection at emit time. +func Once1[T any](e *Emitter, event any, fn func(T)) *Emitter { + e.onces[event] = append(e.onces[event], listenerEntry{ + call: func(args []any) { + if len(args) > 0 { + fn(args[0].(T)) + } + }, + }) + return e +} + +// buildEntry creates a listenerEntry for the given listener. +// Returns (entry, true) on success, (zero, false) if listener is not a Func. +// +// Fast path: common signatures are wrapped with a direct type assertion so +// that no reflection happens when the listener is actually called. +// Slow path: an unknown function type is wrapped with a pooled-reflect +// wrapper that avoids heap allocation for 0–3 argument calls. +// +// 快速路径同样必须记录 ptr:RemoveListener 按 ptr 匹配,漏写会导致 +// 这类监听器永远摘除不掉(实测 recvChannelJoinConfirm 泄漏,每包 +// 数据都重复进入该监听器,日志被 250 行/秒的 DEBUG 刷屏)。 +func buildEntry(listener any) (listenerEntry, bool) { + // Fast path — type-assert well-known signatures. + switch fn := listener.(type) { + case func(): + return listenerEntry{call: func(_ []any) { fn() }, ptr: reflect.ValueOf(listener).Pointer()}, true + case func(error): + return listenerEntry{call: func(args []any) { + var err error + if len(args) > 0 && args[0] != nil { + err = args[0].(error) + } + fn(err) + }, ptr: reflect.ValueOf(listener).Pointer()}, true + case func([]byte): + return listenerEntry{call: func(args []any) { + if len(args) > 0 { + fn(args[0].([]byte)) + } + }, ptr: reflect.ValueOf(listener).Pointer()}, true + case func(uint16): + return listenerEntry{call: func(args []any) { + if len(args) > 0 { + fn(args[0].(uint16)) + } + }, ptr: reflect.ValueOf(listener).Pointer()}, true + } + + // Reflection fallback — verify the value is a Func, then build a wrapper + // that reuses pooled []reflect.Value slices to avoid per-call allocation. + rv := reflect.ValueOf(listener) + if rv.Kind() != reflect.Func { + return listenerEntry{}, false + } + t := rv.Type() + numIn := t.NumIn() + ptr := rv.Pointer() + + // Pre-capture the argument types to avoid repeated Type().In() calls. + ins := make([]reflect.Type, numIn) + for i := range ins { + ins[i] = t.In(i) + } + + call := func(args []any) { + if numIn == 0 { + rv.Call(nil) + return + } + + // Borrow a pre-sized slice from the pool (avoids allocation for 0–3 args). + var vals []reflect.Value + var poolPtr *[]reflect.Value + if numIn < len(rvPool) { + poolPtr = rvPool[numIn].Get().(*[]reflect.Value) + vals = (*poolPtr)[:numIn] + } else { + vals = make([]reflect.Value, numIn) + } + + for i := range numIn { + if i < len(args) && args[i] != nil { + vals[i] = reflect.ValueOf(args[i]) + } else { + vals[i] = reflect.Zero(ins[i]) + } + } + rv.Call(vals) + + // Zero out borrowed slice before returning to pool to avoid memory leaks. + if poolPtr != nil { + for i := range vals { + vals[i] = reflect.Value{} + } + rvPool[numIn].Put(poolPtr) + } + } + + return listenerEntry{call: call, ptr: ptr}, true +} diff --git a/emission/emitter_test.go b/emission/emitter_test.go new file mode 100644 index 0000000..2eba15e --- /dev/null +++ b/emission/emitter_test.go @@ -0,0 +1,46 @@ +package emission + +import "testing" + +// 回归:func([]byte) 快速路径此前不写 ptr,Off/RemoveListener 永远删不掉, +// 并行 MCS 通道加入的持久监听器泄漏后每包数据重复进入(日志刷屏)。 +func TestRemoveListenerFastPathFuncBytes(t *testing.T) { + e := NewEmitter() + calls := 0 + fn := func(s []byte) { calls++ } + e.On("data", fn) + e.Emit("data", []byte("x")) + e.Emit("data", []byte("y")) + if calls != 2 { + t.Fatalf("calls=%d, want 2", calls) + } + e.Off("data", fn) + e.Emit("data", []byte("z")) + if calls != 2 { + t.Fatalf("Off 未生效: calls=%d, want 2", calls) + } +} + +// func(error) 与 func() 快速路径同样必须可摘除。 +func TestRemoveListenerFastPathMisc(t *testing.T) { + e := NewEmitter() + n := 0 + fnErr := func(err error) { n++ } + e.On("error", fnErr) + e.Emit("error", nil) + e.Off("error", fnErr) + e.Emit("error", nil) + if n != 1 { + t.Fatalf("func(error) Off 未生效: n=%d", n) + } + + m := 0 + fn0 := func() { m++ } + e.On("tick", fn0) + e.Emit("tick") + e.Off("tick", fn0) + e.Emit("tick") + if m != 1 { + t.Fatalf("func() Off 未生效: m=%d", m) + } +} diff --git a/go.mod b/go.mod new file mode 100644 index 0000000..e9e7e1f --- /dev/null +++ b/go.mod @@ -0,0 +1,8 @@ +module git.zeroonesoft.cn/golib/rdplib + +go 1.26.3 + +require ( + github.com/lunixbochs/struc v0.0.0-20200707160740-784aaebc1d40 + golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d +) diff --git a/go.sum b/go.sum new file mode 100644 index 0000000..b078551 --- /dev/null +++ b/go.sum @@ -0,0 +1,4 @@ +github.com/lunixbochs/struc v0.0.0-20200707160740-784aaebc1d40 h1:EnfXoSqDfSNJv0VBNqY/88RNnhSGYkrHaO0mmFGbVsc= +github.com/lunixbochs/struc v0.0.0-20200707160740-784aaebc1d40/go.mod h1:vy1vK6wD6j7xX6O6hXe621WabdtNkou2h7uRtTfRMyg= +golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d h1:sK3txAijHtOK88l68nt020reeT1ZdKLIYetKl95FzVY= +golang.org/x/crypto v0.0.0-20220622213112-05595931fe9d/go.mod h1:IxCIyHEi3zRg3s0A5j5BB6A9Jmi73HwBIUl50j+osU4= diff --git a/grdp.go b/grdp.go new file mode 100644 index 0000000..4bb0d14 --- /dev/null +++ b/grdp.go @@ -0,0 +1,2069 @@ +package grdp + +import ( + "errors" + "fmt" + "image" + "log/slog" + "net" + "os" + "runtime/debug" + "strings" + "sync" + "sync/atomic" + "time" + "unsafe" + + "git.zeroonesoft.cn/golib/rdplib/plugin" + "git.zeroonesoft.cn/golib/rdplib/plugin/cliprdr" + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpdr" + "git.zeroonesoft.cn/golib/rdplib/plugin/drdynvc" + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpedisp" + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpgfx" + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpsnd" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" + "git.zeroonesoft.cn/golib/rdplib/protocol/pdu" + "git.zeroonesoft.cn/golib/rdplib/protocol/sec" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/gcc" + "git.zeroonesoft.cn/golib/rdplib/protocol/tpkt" + "git.zeroonesoft.cn/golib/rdplib/protocol/x224" +) + +// mouseCoalescer holds all state for mouse-move coalescing. +// High-frequency move events are collapsed into at most one network PDU per +// mouseCoalesceInterval, with the latest position always winning. +type mouseCoalescer struct { + mu sync.Mutex + pending bool + x, y int + timer *time.Timer + lastTx time.Time + pdu pdu.PointerEvent + pduBuf [1]pdu.InputEventsInterface +} + +// wheelCoalescer holds all state for wheel-scroll coalescing (vertical and +// horizontal). Rapid scroll events are accumulated over mouseCoalesceInterval +// and sent as a single PDU whose rotation value is the sum of all deltas in +// that window. accum/haccum are stored in RDP WHEEL_DELTA units (120 per +// physical notch); haccum is the horizontal axis (positive = scroll right). +type wheelCoalescer struct { + mu sync.Mutex + accum float64 + haccum float64 + timer *time.Timer + lastTx time.Time + pdu pdu.PointerEvent + pduBuf [1]pdu.InputEventsInterface +} + +// stubChannel is a no-op virtual channel handler for channels the server +// expects to be present (e.g. rdpdr, cliprdr) but that we don't process. +type stubChannel struct { + name string + option uint32 + sender core.ChannelSender +} + +func (s *stubChannel) GetType() (string, uint32) { return s.name, s.option } +func (s *stubChannel) Sender(f core.ChannelSender) { s.sender = f } +func (s *stubChannel) Process(data []byte) {} + +type RdpClient struct { + hostPort string // ip:port + width int + height int + kbdLayout uint32 + keyboardType uint32 + keyboardSubType uint32 + timezoneName string + timezoneBias int + tpkt *tpkt.TPKT + x224 *x224.X224 + mcs *t125.MCSClient + sec *sec.Client + pdu *pdu.Client + channels *plugin.Channels + eventReady atomic.Bool + decompressPool sync.Pool // pools []uint8 buffers for bitmap decompression + flipLinePool sync.Pool // pools line-sized []uint8 buffers for bitmap vertical flip + bmpEmitLogged int // 临时诊断:前 10 次位图事件计数 + closed atomic.Bool + + // credentials stored for reconnection + domain string + user string + password string + + // stored callbacks for re-registration on reconnect + onErrorFn func(e error) + onCloseFn func() + onSuccessFn func() + onReadyFn func() + onBitmapPaintFn func([]Bitmap) + onPointerHideFn func() + onPointerCachedFn func(uint16) + onPointerDefaultFn func() + onPointerUpdateFn func(uint16, uint16, uint16, uint16, uint16, uint16, []byte, []byte) + onAudioFn func(rdpsnd.AudioFormat, []byte) + onAudioResetFn func() + onH264RawFn func(destX, destY, w, h int, isKey bool, data []byte, regions []int32) + onH264I420Fn func(destX, destY, w, h int, y []byte, yStride int, u []byte, uStride int, v []byte, vStride int) + onH264NV12Fn func(destX, destY, w, h int, y []byte, yStride int, uv []byte, uvStride int) + onDecoderBrokenFn func() + + // clipboard callbacks and handler + onClipboardFn func(text string) // remote → local + getClipboardFn func() string // local → remote + + onClipboardImageFn func(png []byte) // remote → local(PNG) + onClipboardHTMLFn func(string) // remote → local(HTML Format) + getClipboardHTMLFn func() string // local → remote(HTML Format,空串=无) + getClipboardImgFn func() []byte // local → remote(PNG,无图返回 nil) + + // 文件剪贴板(CF_HDROP + FileContentsRequest/Response,MS-RDPECLIP §3.1.5.4) + onClipboardFilesFn func(names []string) // remote → local:远端剪贴板含文件 + onClipboardFileDataFn func(index int, name string, data []byte) // 单文件下载完成/失败 + onFileProgressFn func(index int, received, total int64) // 下载进度(每段一次) + + // certVerifierFn:TLS 服务器证书 SHA-256 指纹校验(TOFU),nil = 不校验 + certVerifierFn func(sha256Fp []byte) error + + // 位图缓存(stage6 6.4b M1 会话内缓存):CacheBitmapV2 存入, + // MemBlt 按键引用回贴。键 = cacheId<<16 | cacheIndex。 + bitmapCache map[uint32]*Bitmap + bitmapCacheFIFO []uint32 + bmpCacheStores atomic.Uint64 + bmpCacheHits atomic.Uint64 + // gfxCacheStore:GFX/bitmap 两条管线共用的持久缓存桥(6.4/6.4b)。 + gfxCacheStore rdpgfx.GfxCacheStore + + // 连接参数(mstsc 基本选项): + // audioMode: "local"(默认,本机播放)| "none"(不播放)| "remote"(远端播放) + audioMode string + // colorDepth: 0/32 = 默认(WANT_32BPP_SESSION),16/24 生效于传统位图管线 + colorDepth int + // perfFlags/perfFlagsSet: Client Info performanceFlags;未设置保持库默认 + perfFlags uint32 + perfFlagsSet bool + + // rdpsndHandler 在 doLogin 内创建(audioMode != "remote" 时), + // 供同连接周期的 AUDIO_PLAYBACK DVC 适配器复用。 + rdpsndHandler *rdpsnd.Handler + + cliprdrHandler *cliprdr.CliprdrHandler + + // 驱动器重定向(MS-RDPEFS):SetDriveRedirect(true) 后 doLogin 注册 + // 完整 rdpdr 处理器(替换 stub),设备内容由异步 Filesystem 桥提供。 + driveRedirect bool + driveLabel string + driveFS rdpdr.Filesystem // Login 前暂存,处理器创建时补挂 + rdpdrHandler *rdpdr.Handler + + // reconnectMu serialises concurrent Reconnect() calls. + // reconnecting is also set during async server redirects to suppress + // user-facing callbacks while the transport is being re-established. + reconnectMu sync.Mutex + reconnecting atomic.Bool + + // mouse and wheel hold all coalescing state for pointer input. + mouse mouseCoalescer + wheel wheelCoalescer + + // gfxHandler is the active RDPGFX handler; nil when not connected. + // Stored here so closeTransport() can stop its goroutines. + gfxHandler *rdpgfx.GfxHandler + + // avc444Disabled, when true, limits CAPS_ADVERTISE to v8.1 so the server + // uses AVC420 only. Set via DisableAVC444() before Connect(); preserved + // across reconnects. + avc444Disabled bool + + // gfxNoAVC, when true, advertises v10.x caps with AVC_DISABLED so the + // server keeps using the RDPGFX channel with ClearCodec/RFX Progressive + // instead of H.264. Set via NoAVC() before Connect(). + gfxNoAVC bool + + // pduRecorderFn backs SetPduRecorder: the wasm layer may arm recording + // before the client (and its GfxHandler) exists; the callback is attached + // at GfxHandler creation. + pduRecorderFn func(kind byte, codecId uint16, sw, sh, x, y, w, h uint32, payload []byte) + + // rejectGFX, when true, rejects the RDPGFX dynvc channel entirely so the + // server falls back to legacy bitmap updates (compat mode). + rejectGFX bool + + // dispHandler is the active MS-RDPEDISP handler; nil when not connected. + // Used by SetResolution to send MONITOR_LAYOUT PDUs. + dispHandler *rdpedisp.Handler + + dialer func(hostPort string) (net.Conn, error) +} + +const mouseCoalesceInterval = 16 * time.Millisecond + +// Bitmap is a single rendered region delivered to the OnBitmap callback. +// +// Lifecycle: Bitmap.Data is borrowed from an internal buffer pool and is +// only valid for the duration of the synchronous OnBitmap callback. After +// the callback returns the slice may be returned to the pool and overwritten +// by subsequent updates. Callers that need to retain the pixels (e.g. to +// hand them to an asynchronous paint goroutine) MUST copy the bytes before +// the callback returns. +type Bitmap struct { + DestLeft int + DestTop int + DestRight int + DestBottom int + Width int + Height int + BitsPerPixel int + Data []byte +} + +// FillRGBA converts the bitmap's pixel data to RGBA format, writing into dst. +// If dst is nil or has the wrong dimensions a new *image.RGBA is allocated. +// Callers that process tiles of stable dimensions can reuse the same *image.RGBA +// across frames to avoid repeated heap allocations: +// +// var tile *image.RGBA +// tile = bm.FillRGBA(tile) +func (bm *Bitmap) FillRGBA(dst *image.RGBA) *image.RGBA { + if dst == nil || dst.Bounds().Dx() != bm.Width || dst.Bounds().Dy() != bm.Height { + dst = image.NewRGBA(image.Rect(0, 0, bm.Width, bm.Height)) + } + pix := dst.Pix + data := bm.Data + + // Per-format specialised loops avoid a per-pixel switch and let the + // compiler hoist bounds checks and emit tight, branch-free inner code. + switch bm.BitsPerPixel { + case 1: + // 16-bit RGB555 stored big-endian in two bytes. + n := len(pix) >> 2 + if len(data) < n*2 { + n = len(data) / 2 + } + rgb555BatchToRGBA(pix, data, n) + case 2: + // 16-bit RGB565 stored big-endian in two bytes. + n := len(pix) >> 2 + if len(data) < n*2 { + n = len(data) / 2 + } + rgb565BatchToRGBA(pix, data, n) + default: + // 24/32-bit BGR(A) → RGBA with stride = bm.BitsPerPixel. + stride := bm.BitsPerPixel + n := len(pix) >> 2 + if len(data) < n*stride { + n = len(data) / stride + } + if stride == 4 { + // BGRA32 is the common case; use the SIMD-accelerated path. + bgr32BatchToRGBA(pix, data, n) + } else { + // BGR24 (stride==3) and any other depth: scalar fallback. + // Write each pixel as a single 32-bit store to let the compiler + // vectorise the loop (avoids 4 separate byte stores per pixel). + for i := range n { + s := i * stride + *(*uint32)(unsafe.Pointer(&pix[i*4])) = + uint32(data[s+2]) | uint32(data[s+1])<<8 | uint32(data[s])<<16 | 0xFF000000 + } + } + } + return dst +} + +// RGBA converts the bitmap pixel data to an *image.RGBA. +// A new *image.RGBA is allocated on each call. If the caller processes tiles +// of the same dimensions across frames, prefer FillRGBA to avoid allocations. +func (bm *Bitmap) RGBA() *image.RGBA { + return bm.FillRGBA(nil) +} + +// SwapRB swaps the red and blue byte of every 32-bit pixel in p in place and +// forces the alpha byte to 0xFF, converting between RGBA and BGRA byte order +// (the operation is its own inverse). It reuses the SIMD-accelerated batch +// converter (SSE2 on amd64, NEON on arm64), so callers that need BGRA pixels +// for an SDL_PIXELFORMAT_BGRA32 texture can convert an *image.RGBA buffer with +// a single vectorised pass instead of a scalar per-byte swap loop. len(p) must +// be a multiple of 4; any trailing bytes are ignored. +func SwapRB(p []byte) { + bgr32BatchToRGBA(p, p, len(p)/4) +} + +// FillBGRA converts the bitmap pixel data into packed BGRA32 (4 bytes/pixel, +// B at offset 0) and writes it into dst, growing dst if necessary. The caller +// may pass a previously returned slice to reuse the allocation across calls +// (common for tiled rendering where many same-sized bitmaps are converted in a +// loop). The returned slice has length == bm.Width * bm.Height * 4. +// +// Unlike RGBA() + SwapRB, FillBGRA avoids an intermediate *image.RGBA +// allocation and the second-pass R/B swap for BGR24/BGRA32 inputs. For +// BGR24 the output is a single-pass direct pack; for BGRA32 it is a memcopy. +func (bm *Bitmap) FillBGRA(dst []byte) []byte { + n := bm.Width * bm.Height + need := n * 4 + if cap(dst) < need { + dst = make([]byte, need) + } + dst = dst[:need] + data := bm.Data + + switch bm.BitsPerPixel { + case 2: + // 16-bit RGB565: convert to RGBA then swap R↔B in-place → BGRA. + if len(data) < n*2 { + n = len(data) / 2 + } + rgb565BatchToRGBA(dst, data, n) + bgr32BatchToRGBA(dst, dst, n) + default: + stride := bm.BitsPerPixel + if stride < 3 { + stride = 4 // treat unknown as BGRA32 + } + if stride == 4 { + // BGRA32 source — already in the target format; bulk copy. + if len(data) < n*4 { + n = len(data) / 4 + } + copy(dst[:n*4], data[:n*4]) + } else { + // BGR24 source — pack B,G,R directly with A=0xFF; no R/B swap. + if len(data) < n*3 { + n = len(data) / 3 + } + for i := range n { + s := i * 3 + *(*uint32)(unsafe.Pointer(&dst[i*4])) = + uint32(data[s]) | uint32(data[s+1])<<8 | uint32(data[s+2])<<16 | 0xFF000000 + } + } + } + return dst +} + +func NewRdpClient(host string, width, height int, dialer func(string) (net.Conn, error)) *RdpClient { + g := &RdpClient{ + hostPort: host, + width: width, + height: height, + kbdLayout: uint32(gcc.US), + keyboardType: uint32(gcc.KT_IBM_101_102_KEYS), + keyboardSubType: 0, + timezoneName: "UTC", + dialer: dialer, + decompressPool: sync.Pool{ + New: func() any { return []uint8(nil) }, + }, + flipLinePool: sync.Pool{ + New: func() any { return []uint8(nil) }, + }, + } + // Point the cached single-element slices at the cached PDU fields so + // sendMouseMoveLocked / sendWheelLocked need no per-call allocations. + g.mouse.pduBuf[0] = &g.mouse.pdu + g.wheel.pduBuf[0] = &g.wheel.pdu + return g +} + +var keyboardLayoutMap = map[string]uint32{ + "ARABIC": uint32(gcc.ARABIC), + "BULGARIAN": uint32(gcc.BULGARIAN), + "CHINESE_US_KEYBOARD": uint32(gcc.CHINESE_US_KEYBOARD), + "CZECH": uint32(gcc.CZECH), + "DANISH": uint32(gcc.DANISH), + "GERMAN": uint32(gcc.GERMAN), + "GREEK": uint32(gcc.GREEK), + "US": uint32(gcc.US), + "SPANISH": uint32(gcc.SPANISH), + "FINNISH": uint32(gcc.FINNISH), + "FRENCH": uint32(gcc.FRENCH), + "HEBREW": uint32(gcc.HEBREW), + "HUNGARIAN": uint32(gcc.HUNGARIAN), + "ICELANDIC": uint32(gcc.ICELANDIC), + "ITALIAN": uint32(gcc.ITALIAN), + "JAPANESE": uint32(gcc.JAPANESE), + "KOREAN": uint32(gcc.KOREAN), + "DUTCH": uint32(gcc.DUTCH), + "NORWEGIAN": uint32(gcc.NORWEGIAN), +} + +var keyboardTypeMap = map[string]uint32{ + "IBM_PC_XT_83_KEY": uint32(gcc.KT_IBM_PC_XT_83_KEY), + "OLIVETTI": uint32(gcc.KT_OLIVETTI), + "IBM_PC_AT_84_KEY": uint32(gcc.KT_IBM_PC_AT_84_KEY), + "IBM_101_102_KEYS": uint32(gcc.KT_IBM_101_102_KEYS), + "NOKIA_1050": uint32(gcc.KT_NOKIA_1050), + "NOKIA_9140": uint32(gcc.KT_NOKIA_9140), + "JAPANESE": uint32(gcc.KT_JAPANESE), +} + +// SetKeyboardLayout sets the keyboard layout by name (e.g. "US", "FRENCH"). +// Must be called before Login. +func (g *RdpClient) SetKeyboardLayout(layout string) { + if v, ok := keyboardLayoutMap[strings.ToUpper(layout)]; ok { + g.kbdLayout = v + } else { + slog.Warn("Unknown keyboard layout, falling back to US", "layout", layout) + g.kbdLayout = uint32(gcc.US) + } +} + +// SetKeyboardType sets the keyboard type by name (e.g. "IBM_101_102_KEYS"). +// Must be called before Login. +func (g *RdpClient) SetKeyboardType(keyboardType string) { + if v, ok := keyboardTypeMap[strings.ToUpper(keyboardType)]; ok { + g.keyboardType = v + } else { + slog.Warn("Unknown keyboard type, falling back to IBM_101_102_KEYS", "keyboardType", keyboardType) + g.keyboardType = uint32(gcc.KT_IBM_101_102_KEYS) + } +} + +// SetTimezone sets the client timezone reported in the Client Info PDU. +// name is a Windows timezone registry key name (e.g. "UTC", "China Standard +// Time"); biasMinutes = UTC minus local time in minutes (e.g. -480 for UTC+8). +// Must be called before Login. +func (g *RdpClient) SetTimezone(name string, biasMinutes int) { + g.timezoneName = name + g.timezoneBias = biasMinutes +} + +// DisableAVC444 prevents the client from advertising AVC444/AVC444v2 support. +// When called before Login, the RDPGFX CAPS_ADVERTISE is limited to v8.1 +// (AVC420 only), so the server will never send LC=2 chroma-upgrade frames. +// This avoids the colour distortion seen with VirtualBox VRDE, which sends +// LC=2 data but does not include stream2 in LC=0 IDR packets. +// The setting is preserved across automatic reconnects. +func (g *RdpClient) DisableAVC444() *RdpClient { + g.avc444Disabled = true + return g +} + +// RejectGFXChannel rejects the RDPGFX dynamic channel entirely, forcing the +// server to fall back to legacy bitmap updates. Useful against servers whose +// graphics pipeline misbehaves. Must be called before Login. +func (g *RdpClient) RejectGFXChannel() { + g.rejectGFX = true +} + +// NoAVC keeps the RDPGFX channel alive while forbidding H.264: CAPS_ADVERTISE +// includes v10.x sets with RDPGFX_CAPS_FLAG_AVC_DISABLED, so the server encodes +// with ClearCodec + RFX Progressive ("RemoteFX mode"). Must be called before +// Login; preserved across automatic reconnects. +func (g *RdpClient) NoAVC() *RdpClient { + g.gfxNoAVC = true + return g +} + +// SetGfxCacheStore installs the persistent bitmap cache store (MS-RDPEGFX +// CacheImport): SurfaceToCache entries go to the store, and each connection +// offers persisted entries to the server after the caps exchange. Must be +// called after client creation, before the GFX channel finishes negotiation. +func (g *RdpClient) SetGfxCacheStore(s rdpgfx.GfxCacheStore) { + g.gfxCacheStore = s + if h := g.gfxHandler; h != nil { + h.SetPersistentCacheStore(s) + } +} + +// SetPduRecorder installs a callback receiving every wire-to-surface bitmap +// payload for offline replay analysis (garbled-screen debugging harness). +func (g *RdpClient) SetPduRecorder(fn func(kind byte, codecId uint16, surfW, surfH, x, y, w, h uint32, payload []byte)) { + g.pduRecorderFn = fn + if h := g.gfxHandler; h != nil { + h.SetPduRecorder(fn) + } +} + +func bpp(BitsPerPixel uint16) int { + switch BitsPerPixel { + case 15, 16: + return 2 + case 24: + return 3 + case 32: + return 4 + default: + slog.Error("invalid bitmap data format", "BitsPerPixel", BitsPerPixel) + return 0 + } +} + +// mouseButtonFlag returns the PTRFLAGS constant for button index 0/1/2. +func mouseButtonFlag(button int) uint16 { + switch button { + case 0: + return pdu.PTRFLAGS_BUTTON1 + case 2: + return pdu.PTRFLAGS_BUTTON2 + case 1: + return pdu.PTRFLAGS_BUTTON3 + default: + return pdu.PTRFLAGS_MOVE + } +} + +func (g *RdpClient) Login(domain string, user string, password string) error { + slog.Debug("Login", "Host", g.hostPort, "domain", domain, "user", user) + + g.domain = domain + g.user = user + g.password = password + + return g.doLogin(nil) +} + +// doLogin establishes an RDP connection. +// When routingToken is non-nil it replaces the username cookie in the +// x224 Connection Request (required for Server Redirection). +func (g *RdpClient) doLogin(routingToken []byte) error { + conn, err := g.dialer(g.hostPort) + if err != nil { + return fmt.Errorf("[dial err] %v", err) + } + + host, _, _ := net.SplitHostPort(g.hostPort) + socketLayer := core.NewSocketLayer(conn, host) + socketLayer.SetCertVerifier(g.certVerifierFn) + g.tpkt = tpkt.New(socketLayer, nla.NewNTLMv2(g.domain, g.user, g.password)) + g.x224 = x224.New(g.tpkt) + g.mcs = t125.NewMCSClient(g.x224, g.kbdLayout, g.keyboardType, g.keyboardSubType) + g.sec = sec.NewClient(g.mcs) + if g.perfFlagsSet { + g.sec.SetPerformanceFlags(g.perfFlags) + } + // 声音模式的协议层声明(与通道注册行为互补,对齐 mstsc) + switch g.audioMode { + case "none": + g.sec.SetNoAudioPlayback() + case "remote": + g.sec.SetRemoteConsoleAudio() + } + g.pdu = pdu.NewClient(g.sec) + g.channels = plugin.NewChannels(g.sec) + + // Wire user-registered callbacks now that g.pdu is initialised. + // This allows callers to invoke On* methods before Login. + g.reregisterCallbacks() + + // Wire RemoteFX surface decoder so the pdu layer can decode + // codecID=3 in surface bitmap commands without importing rdpgfx. + pdu.DecodeRemoteFX = rdpgfx.DecodeSurfaceRFX + + g.mcs.SetClientDesktop(uint16(g.width), uint16(g.height)) + // 客户端名随机化(RDPDR-2 假设实验):M1 验收成功时 ClientName 为 wasm + // os.Hostname 的固定值 "js";引入随机名后所有会话的 \\tsclient\<共享名> + // 打开均报"试图访问无效的地址"且 rdpdr 零 IRP。ClientName 是成功/失败 + // 之间唯一的客户端侧系统差异——临时禁用随机化以隔离变量。 + if os.Getenv("WEBRDP_RANDOM_CLIENT_NAME") == "1" { + cn := fmt.Sprintf("web%06x", uint64(time.Now().UnixNano())&0xffffff) + g.mcs.SetClientName(cn) + slog.Info("client name", "name", cn) + } + if g.colorDepth != 0 { + g.mcs.SetSessionColorDepth(g.colorDepth) + } + + // Register channels in order: rdpdr, rdpsnd, cliprdr, drdynvc + // (matching the channel order that Windows servers expect) + + // rdpdr (Device Redirection) — stub, required for server to enable audio。 + // 启用驱动器重定向时注册完整处理器(宣告文件系统设备并处理 IO), + // 否则保持 stub。 + if g.driveRedirect { + label := g.driveLabel + if label == "" { + label = "local" + } + // DosName 与 DeviceData 须同名(服务端建 \\tsclient\<共享名> UNC + // 映射的键,FreeRDP 同款);label 同时作卷标。 + // (2026-09-13 凌晨实验结论:DosName 固定 "C" 与 label 同名行为 + // 一致——rdpdr 通道同样在验证序列后被服务端关闭,DosName 已彻底 + // 排除,见 doc/RDPDR-2.md 与 doc/history/stage6-plan.md。) + g.rdpdrHandler = rdpdr.NewHandler(label, label) + if g.driveFS != nil { + g.rdpdrHandler.SetFilesystem(g.driveFS) + } + g.channels.Register(g.rdpdrHandler) + } else { + g.channels.Register(&stubChannel{name: "rdpdr", + option: plugin.CHANNEL_OPTION_INITIALIZED | plugin.CHANNEL_OPTION_ENCRYPT_RDP | plugin.CHANNEL_OPTION_COMPRESS_RDP}) + } + g.mcs.SetClientDeviceRedirection() + + // RDPSND (Audio Output) handler — static virtual channel + DVC paths + // 声音重定向三模式(对应 mstsc 远程音频播放): + // remote(远端播放):不注册 rdpsnd / AUDIO_PLAYBACK 通道, + // 服务器音频走本机扬声器; + // none(不播放):照常注册协商,但 wave 数据丢弃——服务器认为 + // 音频已被重定向而保持静音; + // local(本机播放,默认):注册并播放。 + switch g.audioMode { + case "remote": + // 不注册任何音频通道 + default: + rdpsndHandler := rdpsnd.NewHandler(func(format rdpsnd.AudioFormat, data []byte) { + if g.onAudioFn != nil { + g.onAudioFn(format, data) + } + }) + if g.audioMode == "none" { + rdpsndHandler.SetMuted(true) + } + rdpsndHandler.SetAudioResetCallback(func() { + if g.onAudioResetFn != nil { + g.onAudioResetFn() + } + }) + g.channels.Register(rdpsndHandler) + g.mcs.SetClientSoundProtocol() + + // 音频 DVC 适配器在 doLogin 后半段注册(需要 rdpsndHandler 存活), + // 用局部暂存传递;remote 模式下两通道同样不注册。 + g.rdpsndHandler = rdpsndHandler + } + + // cliprdr (Clipboard) — cross-platform text clipboard handler + cliprdrHandler := cliprdr.NewHandler( + func(text string) { + if g.onClipboardFn != nil { + g.onClipboardFn(text) + } + }, + func() string { + if g.getClipboardFn != nil { + return g.getClipboardFn() + } + return "" + }, + ) + cliprdrHandler.SetImageCallbacks( + func(png []byte) { + if g.onClipboardImageFn != nil { + g.onClipboardImageFn(png) + } + }, + func() []byte { + if g.getClipboardImgFn != nil { + return g.getClipboardImgFn() + } + return nil + }, + ) + cliprdrHandler.SetHTMLCallbacks( + func(html string) { + if g.onClipboardHTMLFn != nil { + g.onClipboardHTMLFn(html) + } + }, + func() string { + if g.getClipboardHTMLFn != nil { + return g.getClipboardHTMLFn() + } + return "" + }, + ) + cliprdrHandler.SetFileCallbacks( + func(names []string) { + if g.onClipboardFilesFn != nil { + g.onClipboardFilesFn(names) + } + }, + func(index int, name string, data []byte) { + if g.onClipboardFileDataFn != nil { + g.onClipboardFileDataFn(index, name, data) + } + }, + func(index int, received, total int64) { + if g.onFileProgressFn != nil { + g.onFileProgressFn(index, received, total) + } + }, + ) + g.cliprdrHandler = cliprdrHandler + g.channels.Register(cliprdrHandler) + g.mcs.SetClientClipboard() + + // drdynvc (Dynamic Virtual Channels) + dvcClient := drdynvc.NewDvcClient() + g.channels.Register(dvcClient) + g.mcs.SetClientDynvcProtocol() + + // RDPGFX (Graphics Pipeline) handler + gfxHandler := rdpgfx.NewGfxHandler(func(updates []rdpgfx.BitmapUpdate) { + if g.onBitmapPaintFn == nil { + return + } + bs := make([]Bitmap, len(updates)) + for i, u := range updates { + bs[i] = Bitmap{ + DestLeft: u.DestLeft, + DestTop: u.DestTop, + DestRight: u.DestRight, + DestBottom: u.DestBottom, + Width: u.Width, + Height: u.Height, + BitsPerPixel: u.Bpp, + Data: u.Data, + } + } + g.onBitmapPaintFn(bs) + }) + gfxHandler.SetDecoderBrokenCallback(func() { + slog.Debug("H.264 decoder broken") + if g.onDecoderBrokenFn != nil { + g.onDecoderBrokenFn() + } + }) + gfxHandler.SetKeyframeRequestFunc(func() { + slog.Debug("H.264: requesting keyframe via force refresh") + if g.pdu != nil { + // SendRefreshRect is silently ignored by Windows servers while + // an H.264 video stream is active. Use the suppress→allow + // toggle (SendForceRefresh) which mstsc/FreeRDP rely on to + // reliably trigger a fresh IDR. See protocol/pdu/pdu.go. + g.pdu.SendForceRefresh(uint16(g.width), uint16(g.height)) + } + }) + if g.onH264RawFn != nil { + gfxHandler.SetH264RawCallback(g.onH264RawFn) + } + if g.onH264I420Fn != nil { + gfxHandler.SetI420Callback(g.onH264I420Fn) + } + if g.onH264NV12Fn != nil { + gfxHandler.SetNV12Callback(g.onH264NV12Fn) + } + if g.avc444Disabled { + gfxHandler.SetAVC444Disabled(true) + } + if g.gfxNoAVC { + gfxHandler.SetAVCDisabled(true) + } + if g.pduRecorderFn != nil { + gfxHandler.SetPduRecorder(g.pduRecorderFn) + } + g.gfxHandler = gfxHandler + // 持久缓存桥:SetGfxCacheStore 在 handler 创建前调用,必须在这里补挂 + // (此前 store 在这里丢失,GFX 持久缓存实际从未生效——服务端 0 条 + // SurfaceToCache 掩盖了这一点)。 + if g.gfxCacheStore != nil && !g.rejectGFX { + gfxHandler.SetPersistentCacheStore(g.gfxCacheStore) + } + // bitmap 管线持久缓存(6.4b M2):把跨会话持有的键注册给 finalize, + // 服务器广告 HOST SUPPORT 时经 PERSISTENT_KEY_LIST 上报。 + if g.rejectGFX && g.gfxCacheStore != nil { + g.pdu.SetPersistentKeyList(g.gfxCacheStore.Keys()) + } + gfxHandler.SetSessionSize(uint16(g.width), uint16(g.height)) + if g.rejectGFX { + dvcClient.RegisterRejectedChannel(rdpgfx.ChannelName) + } else { + dvcClient.RegisterHandler(rdpgfx.ChannelName, gfxHandler) + } + + // RDPEDISP (Display Update Virtual Channel) handler — allows requesting + // a resolution change while connected (MS-RDPEDISP). Pass 0,0 so no + // initial MONITOR_LAYOUT PDU is sent: some servers' graphics pipeline + // fails (ERRINFO_GRAPHICS_SUBSYSTEM_FAILED 0x112F) when the desktop is + // resized during GFX surface setup. Resolution changes go through + // SetResolution() instead. + dispHandler := rdpedisp.NewHandler(0, 0) + g.dispHandler = dispHandler + dvcClient.RegisterHandler(rdpedisp.ChannelName, dispHandler) + + // Reject Video Optimized Remoting (VOR) channels so the server keeps + // sending video through the RDPGFX pipeline which we do handle. + // Without this, the server detects video playback (e.g. YouTube) and + // switches to VOR channels that we don't implement, causing the video + // to freeze while audio continues. + dvcClient.RegisterRejectedChannel("Microsoft::Windows::RDS::Video::Control::v08.01") + dvcClient.RegisterRejectedChannel("Microsoft::Windows::RDS::Video::Data::v08.01") + dvcClient.RegisterRejectedChannel("Microsoft::Windows::RDS::Geometry::v08.01") + + // Register DVC audio handlers for both the lossless and lossy variants. + // gnome-remote-desktop requests AUDIO_PLAYBACK_LOSSY_DVC first; if it is + // rejected, gnome-remote-desktop triggers its SVC fallback path which also + // sets prevent_dvc_initialization=true, silently blocking AUDIO_PLAYBACK_DVC + // as well — leaving the client with no audio at all. + // By accepting both channels with the same rdpsnd handler, format negotiation + // (which only advertises PCM) ensures PCM is used regardless of which channel + // gnome-remote-desktop chooses. + // remote(远端播放)模式不注册,让服务器走本机音频。 + if g.rdpsndHandler != nil { + dvcClient.RegisterHandler("AUDIO_PLAYBACK_DVC", rdpsnd.NewDvcAdapter(g.rdpsndHandler)) + dvcClient.RegisterHandler("AUDIO_PLAYBACK_LOSSY_DVC", rdpsnd.NewDvcAdapter(g.rdpsndHandler)) + } + + g.sec.SetUser(g.user) + g.sec.SetPwd(g.password) + g.sec.SetDomain(g.domain) + // 时区:dynamic DST 键名为空会导致服务器 0x112F 断连;默认发 UTC + g.sec.SetClientTimezone(g.timezoneName, g.timezoneBias) + + g.tpkt.SetFastPathListener(g.sec) + g.sec.SetFastPathListener(g.pdu) + g.sec.SetChannelSender(g.mcs) + g.channels.SetChannelSender(g.sec) + + // Wire fast-path output: pdu → sec → tpkt. This enables the much + // shorter Fast-Path Client Input PDU framing for mouse/keyboard events + // (MS-RDPBCGR §2.2.8.1.2). Use is gated at runtime both by capability + // negotiation in the PDU layer and by sec.SendFastPath itself, which + // refuses when legacy RDP encryption is in effect. + g.sec.SetFastPathSender(g.tpkt) + g.pdu.SetFastPathSender(g.sec) + + g.x224.SetRequestedProtocol(x224.PROTOCOL_SSL | x224.PROTOCOL_HYBRID) + if routingToken != nil { + g.x224.SetRoutingToken(routingToken) + } else { + g.x224.SetUsername(g.user) + } + + err = g.x224.Connect() + if err != nil { + return fmt.Errorf("[x224 connect err] %v", err) + } + + // Wait for the RDP handshake to complete or fail. + // Events arrive asynchronously from the TPKT read goroutine. + type connResult struct { + err error + redirect *pdu.ServerRedirectionPDU + } + + ch := make(chan connResult, 4) + send := func(r connResult) { + select { + case ch <- r: + default: + } + } + + // readyFired is set by the "ready" callback. All emitter callbacks + // run synchronously on the TPKT read goroutine, so no mutex needed. + readyFired := false + + g.pdu.On("ready", func() { + g.eventReady.Store(true) + readyFired = true + send(connResult{}) + }) + + g.pdu.On("error", func(err error) { + if !readyFired { + send(connResult{err: err}) + } else { + // Mid-session error: stop accepting input so we don't + // try to write to the now-dead transport. + g.eventReady.Store(false) + } + }) + + // Redirect may arrive before or after "ready". + // Before ready: send to channel for synchronous handling. + // After ready: launch async goroutine (GNOME Remote Desktop + // sends redirect ~5s after the GFX retry's "ready"). + g.pdu.Once("redirect", func(redir *pdu.ServerRedirectionPDU) { + if !readyFired { + send(connResult{redirect: redir}) + } else { + go g.handleRedirect(redir) + } + }) + + // DeactivateAllPDU during an active session means the server is + // reactivating (e.g. desktop resize). Pause input until "ready" + // fires again after the reactivation handshake completes. + g.pdu.On("deactivateAll", func() { + g.eventReady.Store(false) + }) + + select { + case r := <-ch: + if r.err != nil { + g.tpkt.Close() + return fmt.Errorf("[connection err] %v", r.err) + } + if r.redirect != nil { + slog.Debug("Server redirect", "loadBalanceInfo", string(r.redirect.LoadBalanceInfo)) + g.tpkt.Close() + g.eventReady.Store(false) + return g.doLogin(r.redirect.LoadBalanceInfo) + } + // "ready" received — session established. + return nil + case <-time.After(30 * time.Second): + g.tpkt.Close() + return fmt.Errorf("[connection timeout]") + } +} + +// handleRedirect handles a Server Redirection PDU that arrives after +// "ready" (e.g. GNOME Remote Desktop). Runs asynchronously. +func (g *RdpClient) handleRedirect(redir *pdu.ServerRedirectionPDU) { + slog.Debug("Async server redirect", "loadBalanceInfo", string(redir.LoadBalanceInfo)) + g.reconnecting.Store(true) + g.tpkt.Close() + g.eventReady.Store(false) + + err := g.doLogin(redir.LoadBalanceInfo) + g.reconnecting.Store(false) + if err != nil { + slog.Error("handleRedirect: login failed", "err", err) + if g.onErrorFn != nil { + g.onErrorFn(err) + } + return + } + g.reregisterCallbacks() +} + +func (g *RdpClient) Width() int { + return g.width +} + +func (g *RdpClient) Height() int { + return g.height +} + +func (g *RdpClient) OnError(f func(e error)) *RdpClient { + g.onErrorFn = f + if g.pdu != nil { + g.pdu.On("error", func(e error) { + if !g.reconnecting.Load() { + f(e) + } + }) + } + return g +} + +func (g *RdpClient) OnClose(f func()) *RdpClient { + g.onCloseFn = f + if g.pdu != nil { + g.pdu.On("close", func() { + if !g.reconnecting.Load() { + f() + } + }) + } + return g +} + +func (g *RdpClient) OnSuccess(f func()) *RdpClient { + g.onSuccessFn = f + if g.sec != nil { + g.sec.On("success", f) + } + return g +} + +func (g *RdpClient) OnReady(f func()) *RdpClient { + g.onReadyFn = f + if g.pdu != nil { + g.pdu.On("ready", f) + } + return g +} + +// OnBitmap registers a callback for bitmap update events. +// For compressed bitmaps, Bitmap.Data is borrowed from an internal pool and +// is valid only for the duration of the paint call. If you need to retain +// the raw pixel data beyond paint, copy it or call bm.RGBA() inside paint. +func (g *RdpClient) OnBitmap(paint func([]Bitmap)) *RdpClient { + g.onBitmapPaintFn = paint + if g.pdu == nil { + return g + } + g.pdu.On("bitmap", func(rectangles []pdu.BitmapData) { + // 16/24bpp 位图路径曾发生 panic 直接杀死整个 wasm 程序; + // recover 保证连接存活并留下完整堆栈用于定位。 + defer func() { + if r := recover(); r != nil { + slog.Error("bitmap update panic", "err", r, "stack", string(debug.Stack())) + } + }() + bs := make([]Bitmap, 0, len(rectangles)) + var pooled [][]uint8 // track buffers borrowed from pool + + for idx, v := range rectangles { + data := v.BitmapDataStream + wireBpp := v.BitsPerPixel + if wireBpp == 0 { + // Win10 在低色深会话(经 postBeta2ColorDepth 协商)里把 + // bitsPerPixel 置 0:按会话色深处理 + wireBpp = uint16(g.colorDepth) + if wireBpp == 0 { + wireBpp = 32 + } + } + Bpp := bpp(wireBpp) + if Bpp == 0 { + slog.Error("bitmap rect with invalid bpp", + "idx", idx, "count", len(rectangles), + "rect", fmt.Sprintf("%+v", v), + "allRects", fmt.Sprintf("%+v", rectangles)) + continue + } + + if v.Flags&pdu.BITMAP_NO_PROCESSING != 0 { + // Surface command: data is already decoded top-down BGRA + } else if v.IsCompress() { + buf := g.decompressPool.Get().([]uint8) + var ok bool + buf, ok = core.DecompressInto(v.BitmapDataStream, buf, int(v.Width), int(v.Height), Bpp) + if !ok { + // 解码失败的矩形绝不能上屏:缓冲里是半解码+池化残留, + // 画出来就是噪块。跳过并等服务器后续更新修复该区域。 + g.decompressPool.Put(buf) + slog.Warn("skip undecodable bitmap rect", + "dx", v.DestLeft, "dy", v.DestTop, + "dr", v.DestRight, "db", v.DestBottom, + "w", v.Width, "h", v.Height, "bpp", Bpp, + "flags", fmt.Sprintf("0x%04X", v.Flags), + "len", len(v.BitmapDataStream)) + continue + } + data = buf + pooled = append(pooled, buf) + } else { + // Uncompressed bitmaps are bottom-up; flip to top-down. + stride := int(v.Width) * Bpp + h := int(v.Height) + tmp := g.flipLinePool.Get().([]byte) + if cap(tmp) < stride { + tmp = make([]byte, stride) + } else { + tmp = tmp[:stride] + } + for y := 0; y < h/2; y++ { + top := y * stride + bot := (h - 1 - y) * stride + copy(tmp, data[top:top+stride]) + copy(data[top:top+stride], data[bot:bot+stride]) + copy(data[bot:bot+stride], tmp) + } + g.flipLinePool.Put(tmp[:cap(tmp)]) + } + + b := Bitmap{int(v.DestLeft), int(v.DestTop), int(v.DestRight), int(v.DestBottom), + int(v.Width), int(v.Height), Bpp, data} + bs = append(bs, b) + } + if g.bmpEmitLogged < 10 { + g.bmpEmitLogged++ + slog.Warn("BITMAP_EMIT", "rects", len(bs)) + } + paint(bs) + + for _, buf := range pooled { + g.decompressPool.Put(buf[:cap(buf)]) + } + }) + + // 位图缓存(stage6 6.4b M1):CacheBitmapV2 存入,MemBlt 引用回贴。 + // 订单与位图更新同为 fast-path PDU,按到达顺序同步处理,服务器保证 + // MemBlt 先于对应 CacheBitmapV2 的乱序不存在。 + g.bitmapCache = make(map[uint32]*Bitmap) + ordersSeen := 0 + g.pdu.On("orders", func(orderPdus []pdu.OrderPdu) { + ordersSeen++ + if ordersSeen <= 5 { + kinds := make(map[string]int) + for i := range orderPdus { + o := &orderPdus[i] + switch { + case o.CacheBitmapV2 != nil: + kinds["CacheBitmapV2"]++ + case o.Primary != nil && o.Primary.Data != nil: + kinds[fmt.Sprintf("%T", o.Primary.Data)]++ + default: + kinds["empty"]++ + } + } + slog.Warn("ORDERS event", "n", ordersSeen, "pdus", len(orderPdus), "kinds", fmt.Sprintf("%v", kinds)) + } + defer func() { + if r := recover(); r != nil { + slog.Error("orders update panic", "err", r, "stack", string(debug.Stack())) + } + }() + for i := range orderPdus { + o := &orderPdus[i] + if cb := o.CacheBitmapV2; cb != nil { + g.storeBitmapCacheV2(cb) + continue + } + if o.Primary != nil { + if mb, ok := o.Primary.Data.(*pdu.Memblt); ok { + g.drawMemblt(mb) + } + } + } + }) + return g +} + +func (g *RdpClient) OnPointerHide(f func()) *RdpClient { + g.onPointerHideFn = f + if g.pdu != nil { + g.pdu.On("pointer_hide", f) + } + return g +} + +func (g *RdpClient) OnPointerCached(f func(uint16)) *RdpClient { + g.onPointerCachedFn = f + if g.pdu != nil { + g.pdu.On("pointer_cached", f) + } + return g +} + +func (g *RdpClient) OnPointerDefault(f func()) *RdpClient { + g.onPointerDefaultFn = f + if g.pdu != nil { + g.pdu.On("pointer_default", f) + } + return g +} + +// bitmapCacheMaxEntries 单元总量上限(超出按 FIFO 逐出,防御服务器 +// 引用已逐出条目时缓存无限增长)。 +const bitmapCacheMaxEntries = 4096 + +// storeBitmapCacheV2 解压 CacheBitmapV2 次级订单的位图并按 +// cacheId<<16|cacheIndex 存入会话内缓存(DO_NOT_CACHE 条目不存)。 +func (g *RdpClient) storeBitmapCacheV2(cb *pdu.CacheBitmapV2Order) { + if cb.CacheIndex == 0x7FFF || cb.BitmapWidth == 0 || cb.BitmapHeight == 0 { + return + } + // 持久键(6.4b M2):服务器在 CacheBitmapV2 上携带跨会话键。 + // 带键且零数据 = 服务器指示客户端从持久库回填该单元。 + persistentKey := uint64(0) + if cb.Flags&pdu.CBR2_PERSISTENT_KEY_PRESENT != 0 { + persistentKey = uint64(cb.Key2)<<32 | uint64(cb.Key1) + } + Bpp := bpp(uint16(cb.BitmapBpp)) + if Bpp == 0 { + return + } + w, h := int(cb.BitmapWidth), int(cb.BitmapHeight) + + if cb.BitmapLength == 0 && persistentKey != 0 { + if g.gfxCacheStore == nil { + return + } + e, ok := g.gfxCacheStore.Get(persistentKey) + if !ok || e.Width != w || e.Height != h || e.Bpp != uint16(Bpp) || + len(e.Data) != w*h*int(Bpp) { + slog.Debug("bmpcache: persistent backfill miss", "key", persistentKey, + "cacheId", cb.CacheId, "idx", cb.CacheIndex, "w", w, "h", h) + return + } + entry := &Bitmap{Width: w, Height: h, BitsPerPixel: Bpp, Data: e.Data} + g.bitmapCachePut(uint32(cb.CacheId)<<16|uint32(cb.CacheIndex), entry) + return + } + + stride := w * Bpp + var data []byte + var ok bool + if cb.Compressed { + buf := g.decompressPool.Get().([]byte) + var out []byte + out, ok = core.DecompressInto(cb.BitmapDataStream, buf, w, h, Bpp) + if !ok { + g.decompressPool.Put(buf) + slog.Debug("bmpcache: skip undecodable cache bitmap", + "cacheId", cb.CacheId, "idx", cb.CacheIndex, + "w", w, "h", h, "bpp", Bpp, "len", len(cb.BitmapDataStream)) + return + } + data = out + } else { + data = append([]byte(nil), cb.BitmapDataStream...) + } + // RLE 位图自底向上,翻转为自顶向下(与 bitmap 更新路径一致) + mirrorRows(data, stride, h) + + entry := &Bitmap{Width: w, Height: h, BitsPerPixel: Bpp, Data: data} + g.bitmapCachePut(uint32(cb.CacheId)<<16|uint32(cb.CacheIndex), entry) + + // 带持久键的条目跨会话持久化(像素副本,键 = key2<<32|key1) + if persistentKey != 0 && g.gfxCacheStore != nil { + g.gfxCacheStore.Persist(persistentKey, w, h, uint16(Bpp), data) + } + + stores := g.bmpCacheStores.Add(1) + if stores == 1 || stores%500 == 0 { + slog.Info("bmpcache: stored", "stores", stores, "total", len(g.bitmapCache), + "cacheId", cb.CacheId, "idx", cb.CacheIndex, "w", w, "h", h, "bpp", Bpp) + } +} + +// bitmapCachePut 写入会话内缓存单元(cacheId<<16|cacheIndex → 位图), +// 超过 bitmapCacheMaxEntries 按插入序 FIFO 逐出。 +func (g *RdpClient) bitmapCachePut(key uint32, entry *Bitmap) { + if _, exists := g.bitmapCache[key]; !exists { + g.bitmapCacheFIFO = append(g.bitmapCacheFIFO, key) + if len(g.bitmapCacheFIFO) > bitmapCacheMaxEntries { + evict := g.bitmapCacheFIFO[0] + g.bitmapCacheFIFO = g.bitmapCacheFIFO[1:] + delete(g.bitmapCache, evict) + } + } + g.bitmapCache[key] = entry +} + +// drawMemblt 处理 MemBlt 主订单:从缓存取位图按目标坐标回贴。 +// 缓存缺失时跳过(保持既有画面,等服务器后续更新修复该区域)。 +func (g *RdpClient) drawMemblt(mb *pdu.Memblt) { + if g.onBitmapPaintFn == nil { + return + } + key := uint32(mb.CacheId)<<16 | uint32(mb.CacheIdx) + b, ok := g.bitmapCache[key] + if !ok || b == nil { + slog.Debug("bmpcache: memblt cache miss", "cacheId", mb.CacheId, "idx", mb.CacheIdx) + return + } + out := *b + out.DestLeft, out.DestTop = int(mb.X), int(mb.Y) + out.DestRight, out.DestBottom = int(mb.X)+int(mb.Cx), int(mb.Y)+int(mb.Cy) + hits := g.bmpCacheHits.Add(1) + if hits == 1 || hits%500 == 0 { + slog.Info("bmpcache: hit", "hits", hits, "cacheId", mb.CacheId, + "idx", mb.CacheIdx, "x", mb.X, "y", mb.Y, "w", mb.Cx, "h", mb.Cy) + } + g.onBitmapPaintFn([]Bitmap{out}) +} + +// mirrorRows 将 stride 对齐的行序图像原地垂直镜像(bottom-up ↔ top-down)。 +func mirrorRows(buf []byte, stride, height int) { + tmp := make([]byte, stride) + for y := 0; y < height/2; y++ { + top := y * stride + bot := (height - 1 - y) * stride + if top+stride <= len(buf) && bot+stride <= len(buf) { + copy(tmp, buf[top:top+stride]) + copy(buf[top:top+stride], buf[bot:bot+stride]) + copy(buf[bot:bot+stride], tmp) + } + } +} + +func (g *RdpClient) OnPointerUpdate(f func(uint16, uint16, uint16, uint16, uint16, uint16, []byte, []byte)) *RdpClient { + g.onPointerUpdateFn = f + if g.pdu != nil { + g.pdu.On("pointer_update", func(p *pdu.FastPathUpdatePointerPDU) { + w := int(p.Width) + h := int(p.Height) + + // xorBpp 由线路直接携带(TS_POINTER_NEW 首字段,FreeRDP + // update_read_pointer_new 校验 1≤xorBpp≤32)。 + xorBpp := int(p.XorBpp) + if xorBpp == 0 { + xorBpp = 1 + } + slog.Debug("OnPointerUpdate", "cacheIdx", p.CacheIdx, "xorBpp", p.XorBpp, + "hotX", p.HotX, "hotY", p.HotY, "w", p.Width, "h", p.Height, + "andLen", p.MaskLen, "xorLen", p.XorLen) + + // 掩码行序:按 MS-RDPBCGR/FreeRDP(vFlip = xorBpp != 1),彩色 + // 指针(24/32bpp)自底向上存储,翻转为自顶向下;1bpp 单色指针 + // 本身自顶向下。stride 均按 2 字节对齐。 + flip := xorBpp != 1 + xorStride := ((w*xorBpp + 15) / 16) * 2 + andStride := ((w + 15) / 16) * 2 + var xorData []byte + if len(p.Data) > 0 && h > 0 && w > 0 { + xorData = make([]byte, len(p.Data)) + copy(xorData, p.Data) + if flip { + mirrorRows(xorData, xorStride, h) + } + } else { + xorData = p.Data + } + + var andMask []byte + if len(p.Mask) > 0 && h > 0 && w > 0 { + andMask = make([]byte, len(p.Mask)) + copy(andMask, p.Mask) + if flip { + mirrorRows(andMask, andStride, h) + } + } else { + andMask = p.Mask + } + + f(p.CacheIdx, uint16(xorBpp), p.HotX, p.HotY, p.Width, p.Height, andMask, xorData) + }) + } + return g +} + +// OnAudio registers a callback for server audio data. +// The callback receives the AudioFormat describing the PCM data and the raw audio bytes. +// Must be called before Login. +func (g *RdpClient) OnAudio(f func(rdpsnd.AudioFormat, []byte)) *RdpClient { + g.onAudioFn = f + return g +} + +// OnAudioReset registers a callback that is called when the server closes the +// audio channel (e.g. media seek or stream restart). The application should +// flush its audio playback buffer so that stale audio does not keep playing. +// Must be called before Login. +func (g *RdpClient) OnAudioReset(f func()) *RdpClient { + g.onAudioResetFn = f + return g +} + +// OnH264Raw registers a callback that receives raw H.264 NAL unit data when +// the built-in decoder is unavailable (e.g. WASM builds without CGo). +// destX, destY are the top-left canvas coordinates; isKey flags an IDR frame. +// regions 为扁平 [l,t,r,b,...](帧内坐标,右下开区间)的脏矩形:服务器只 +// 保证区域内像素有效,绘制端必须只上屏区域;空切片表示整帧有效。 +// The caller owns data and may retain it beyond the callback. +func (g *RdpClient) OnH264Raw(fn func(destX, destY, w, h int, isKey bool, data []byte, regions []int32)) *RdpClient { + g.onH264RawFn = fn + return g +} + +// OnH264I420 registers a callback that receives decoded H.264 frames in I420 +// planar format (Y, U, V planes with associated strides). When set, the +// decoded frame is NOT delivered via OnBitmap; the caller is responsible for +// rendering it directly (e.g. via an SDL2 IYUV texture for GPU-accelerated +// YUV→RGB conversion). When I420 extraction is unavailable for a frame +// (e.g. non-YUV420P/NV12 formats), grdp falls back to OnBitmap delivery. +// destX, destY are top-left canvas coordinates; w, h are frame dimensions. +// The plane slices are only valid for the duration of the callback; copy them +// if they need to be retained beyond the callback's return. +func (g *RdpClient) OnH264I420(fn func(destX, destY, w, h int, y []byte, yStride int, u []byte, uStride int, v []byte, vStride int)) *RdpClient { + g.onH264I420Fn = fn + if g.gfxHandler != nil { + g.gfxHandler.SetI420Callback(fn) + } + return g +} + +// OnH264NV12 registers a callback that receives decoded H.264 frames in NV12 +// format (Y plane plus interleaved UV plane). This is the fastest SDL2 path +// on platforms whose hardware decoder already outputs NV12 (notably macOS +// VideoToolbox), because callers can upload the planes directly with an NV12 +// texture and avoid NV12->I420 deinterleaving in grdp. When NV12 extraction +// is unavailable for a frame, grdp falls back to OnBitmap delivery. +// destX, destY are top-left canvas coordinates; w, h are frame dimensions. +// The plane slices are only valid for the duration of the callback; copy them +// if they need to be retained beyond the callback's return. +func (g *RdpClient) OnH264NV12(fn func(destX, destY, w, h int, y []byte, yStride int, uv []byte, uvStride int)) *RdpClient { + g.onH264NV12Fn = fn + if g.gfxHandler != nil { + g.gfxHandler.SetNV12Callback(fn) + } + return g +} + +// OnDecoderBroken registers a callback that is invoked when the H.264 decoder +// enters an unrecoverable state (all hard-reset attempts exhausted). When +// this callback is set, grdp does NOT automatically call Reconnect; the +// application is responsible for deciding when to reconnect (e.g. via its +// own stall watchdog). If no callback is registered, grdp falls back to +// the previous behaviour of reconnecting immediately. +func (g *RdpClient) OnDecoderBroken(f func()) *RdpClient { + g.onDecoderBrokenFn = f + return g +} + +// OnClipboard registers callbacks for bidirectional clipboard sharing. +// +// - onRemote is called with the text when the RDP server's clipboard +// content is received (server → client). +// - getLocal is called to retrieve the current local clipboard text +// when the server requests it (client → server). +// +// Must be called before Login. +func (g *RdpClient) OnClipboard(onRemote func(text string), getLocal func() string) *RdpClient { + g.onClipboardFn = onRemote + g.getClipboardFn = getLocal + return g +} + +// OnClipboardImage registers the callback invoked with PNG-encoded bytes when +// the RDP server's clipboard image is received (server → client). Must be +// called before Login. +func (g *RdpClient) OnClipboardImage(onRemoteImage func(png []byte)) *RdpClient { + g.onClipboardImageFn = onRemoteImage + return g +} + +// SetClipboardImageProvider registers the provider used to answer server +// requests for the local clipboard image (client → server). The provider +// returns PNG-encoded bytes, or nil when no image is on the local clipboard. +// Must be called before Login. +func (g *RdpClient) SetClipboardImageProvider(getImage func() []byte) *RdpClient { + g.getClipboardImgFn = getImage + return g +} + +// OnClipboardHTML registers the callback invoked when the remote clipboard +// HTML content (HTML Format, fragment already extracted) arrives. +func (g *RdpClient) OnClipboardHTML(onRemoteHTML func(html string)) *RdpClient { + g.onClipboardHTMLFn = onRemoteHTML + return g +} + +// SetClipboardHTMLProvider registers the provider used to answer server +// requests for the local clipboard HTML (HTML Format). The provider returns +// raw HTML (no CF_HTML envelope), or "" when no HTML is on the local +// clipboard. Must be called before Login. +func (g *RdpClient) SetClipboardHTMLProvider(getHTML func() string) *RdpClient { + g.getClipboardHTMLFn = getHTML + return g +} + +// OnClipboardFiles registers the callback invoked when the remote clipboard +// holds files (CF_HDROP path list received, server → client). Files are NOT +// fetched automatically; call RequestRemoteFile for each file to download. +func (g *RdpClient) OnClipboardFiles(onRemoteFiles func(names []string)) *RdpClient { + g.onClipboardFilesFn = onRemoteFiles + return g +} + +// OnClipboardFileData registers the callback invoked when a file requested +// via RequestRemoteFile has been fully received. data is nil on failure. +func (g *RdpClient) OnClipboardFileData(fn func(index int, name string, data []byte)) *RdpClient { + g.onClipboardFileDataFn = fn + return g +} + +// OnFileTransferProgress registers the callback invoked after every received +// chunk of an in-flight RequestRemoteFile transfer. +func (g *RdpClient) OnFileTransferProgress(fn func(index int, received, total int64)) *RdpClient { + g.onFileProgressFn = fn + return g +} + +// SetLocalFiles stages files as the local clipboard file content (client → +// server). The server will see a CF_HDROP format and fetch bytes on paste. +func (g *RdpClient) SetLocalFiles(files []cliprdr.LocalFile) { + if g.cliprdrHandler != nil { + g.cliprdrHandler.SetLocalFiles(files) + } +} + +// ClearLocalFiles removes locally staged clipboard files. +func (g *RdpClient) ClearLocalFiles() { + if g.cliprdrHandler != nil { + g.cliprdrHandler.ClearLocalFiles() + } +} + +// RequestRemoteFile starts downloading file `index` from the remote clipboard +// file list; progress and completion arrive via the registered callbacks. +func (g *RdpClient) RequestRemoteFile(index int) error { + if g.cliprdrHandler == nil { + return errors.New("client not connected") + } + return g.cliprdrHandler.RequestRemoteFile(index) +} + +// SetCertVerifier registers a TOFU (trust-on-first-use) verifier for the +// server's TLS certificate. The callback receives the SHA-256 of the leaf +// certificate DER; return an error to abort the connection. Must be called +// before Login. +func (g *RdpClient) SetCertVerifier(fn func(sha256Fp []byte) error) *RdpClient { + g.certVerifierFn = fn + return g +} + +// NotifyClipboardChanged tells the server that the local clipboard has +// changed. The UI should call this when it detects a system clipboard +// change (e.g. via polling or a platform clipboard-change signal). +func (g *RdpClient) NotifyClipboardChanged() { + if g.cliprdrHandler != nil { + g.cliprdrHandler.OnLocalClipboardChanged() + } +} + +func (g *RdpClient) notifyGfxLocalInput() { + if gfx := g.gfxHandler; gfx != nil { + gfx.NotifyLocalInput() + } +} + +// newScancodeEvent builds a TS_SCANCODE_EVENT from the 0xE0xx convention +// used throughout this codebase: the E0 prefix must travel as +// KBDFLAGS_EXTENDED with the 8-bit make code in KeyCode — the slow-path +// serializer sends KeyCode verbatim, and a raw 0xE0xx value there is an +// invalid scancode the server silently drops (observed: Delete did +// nothing on the remote). +func newScancodeEvent(sc int, release bool) *pdu.ScancodeKeyEvent { + p := &pdu.ScancodeKeyEvent{} + if sc&0xFF00 == 0xE000 { + p.KeyboardFlags |= pdu.KBDFLAGS_EXTENDED + sc &= 0xFF + } + p.KeyCode = uint16(sc) + if release { + p.KeyboardFlags |= pdu.KBDFLAGS_RELEASE + } + return p +} + +func (g *RdpClient) KeyUp(sc int) { + if !g.eventReady.Load() { + return + } + slog.Debug("KeyUp", "sc", sc) + g.flushMouseMove() + g.flushWheel() + + p := newScancodeEvent(sc, true) + g.pdu.SendInputEvents(pdu.INPUT_EVENT_SCANCODE, []pdu.InputEventsInterface{p}) + g.notifyGfxLocalInput() +} + +// SendUnicodeText sends text as RDP Unicode input events (MS-RDPBCGR +// RDP_INPUT_UNICODE, TS_UNICODE_EVENT). Used for IME-committed text and +// clipboard paste, which have no meaningful scancode representation. +// Each character is sent as a press+release pair; characters outside the +// BMP (no UTF-16 code unit mapping) are skipped. +func (g *RdpClient) SendUnicodeText(text string) { + if !g.eventReady.Load() { + return + } + g.flushMouseMove() + g.flushWheel() + + const maxEventsPerPDU = 14 // fast-path input PDU hard limit is 15 events + events := make([]pdu.InputEventsInterface, 0, maxEventsPerPDU) + flush := func() { + if len(events) == 0 { + return + } + g.pdu.SendInputEvents(pdu.INPUT_EVENT_UNICODE, events) + events = events[:0] + } + for _, r := range text { + if r == '\r' || r == '\n' { + // CR/LF 走 Unicode 事件会被服务端当不可打印字符丢弃(实测 cmd + // 收不到回车,粘贴多行文本/IME 提交带回车的文本无法执行),转成 + // Enter 扫描码按下+释放。扫描码必须走 INPUT_EVENT_SCANCODE PDU: + // 与 Unicode 事件混在同一 PDU 里服务端会按 TS_UNICODE_EVENT + // 解析出控制字符并丢弃(实测 Enter 静默丢失)。 + flush() + g.pdu.SendInputEvents(pdu.INPUT_EVENT_SCANCODE, []pdu.InputEventsInterface{ + newScancodeEvent(0x001C, false), + newScancodeEvent(0x001C, true), + }) + continue + } + if r > 0xFFFF || (r >= 0xD800 && r <= 0xDFFF) { + continue + } + u := uint16(r) + events = append(events, + &pdu.UnicodeKeyEvent{Unicode: u}, + &pdu.UnicodeKeyEvent{Unicode: u, KeyboardFlags: pdu.KBDFLAGS_RELEASE}, + ) + if len(events) >= maxEventsPerPDU { + flush() + } + } + flush() + g.notifyGfxLocalInput() +} + +func (g *RdpClient) KeyDown(sc int) { + if !g.eventReady.Load() { + return + } + slog.Debug("KeyDown", "sc", sc) + g.flushMouseMove() + g.flushWheel() + + p := newScancodeEvent(sc, false) + g.pdu.SendInputEvents(pdu.INPUT_EVENT_SCANCODE, []pdu.InputEventsInterface{p}) + g.notifyGfxLocalInput() +} + +// MouseMove queues a mouse-move event. Successive moves within +// mouseCoalesceInterval are collapsed: only the latest (x,y) is sent. The +// first move in a burst is sent immediately so the server sees no extra +// latency for a single isolated motion. +func (g *RdpClient) MouseMove(x, y int) { + if !g.eventReady.Load() { + return + } + + g.mouse.mu.Lock() + g.mouse.x = x + g.mouse.y = y + g.mouse.pending = true + + now := time.Now() + since := now.Sub(g.mouse.lastTx) + if since >= mouseCoalesceInterval { + // Throttle window has elapsed — send right away. + g.sendMouseMoveLocked(now) + g.mouse.mu.Unlock() + return + } + + // Within throttle window: schedule a flush for the remainder of it + // (unless one is already scheduled). + if g.mouse.timer == nil { + delay := mouseCoalesceInterval - since + g.mouse.timer = time.AfterFunc(delay, g.flushMouseMoveTimer) + } + g.mouse.mu.Unlock() +} + +// flushMouseMove sends any pending mouse-move event synchronously. Called +// before any non-move input event to preserve server-side ordering. +func (g *RdpClient) flushMouseMove() { + g.mouse.mu.Lock() + if g.mouse.timer != nil { + g.mouse.timer.Stop() + g.mouse.timer = nil + } + if g.mouse.pending { + g.sendMouseMoveLocked(time.Now()) + } + g.mouse.mu.Unlock() +} + +// flushMouseMoveTimer is the time.AfterFunc callback. Acquires the lock +// itself and sends whatever's pending. +func (g *RdpClient) flushMouseMoveTimer() { + g.mouse.mu.Lock() + g.mouse.timer = nil + if g.mouse.pending && g.eventReady.Load() { + g.sendMouseMoveLocked(time.Now()) + } + g.mouse.mu.Unlock() +} + +// sendMouseMoveLocked must be called with mouse.mu held. +func (g *RdpClient) sendMouseMoveLocked(now time.Time) { + g.mouse.pdu.PointerFlags = pdu.PTRFLAGS_MOVE + g.mouse.pdu.XPos = uint16(g.mouse.x) + g.mouse.pdu.YPos = uint16(g.mouse.y) + g.mouse.pending = false + g.mouse.lastTx = now + g.pdu.SendInputEvents(pdu.INPUT_EVENT_MOUSE, g.mouse.pduBuf[:]) +} + +// MouseWheel sends a vertical scroll event to the remote desktop. +// delta is the rotation amount in physical notches (1.0 = one click of a +// scroll wheel = Windows WHEEL_DELTA). Fractional values are accepted for +// smooth / high-resolution input devices such as trackpads. +// Positive values scroll up (away from the user); negative values scroll down. +func (g *RdpClient) MouseWheel(delta float64) { + if !g.eventReady.Load() { + return + } + slog.Debug("MouseWheel", "delta", delta) + g.flushMouseMove() + + // Convert notch count to RDP WHEEL_DELTA units (120 per notch). + const wheelDelta = 120 + g.wheel.mu.Lock() + g.wheel.accum += delta * wheelDelta + if g.wheel.accum == 0 { + // Opposite deltas cancelled out; nothing to send. + g.wheel.mu.Unlock() + return + } + + now := time.Now() + since := now.Sub(g.wheel.lastTx) + if since >= mouseCoalesceInterval { + g.sendWheelLocked(now) + g.wheel.mu.Unlock() + return + } + + if g.wheel.timer == nil { + delay := mouseCoalesceInterval - since + g.wheel.timer = time.AfterFunc(delay, g.flushWheelTimer) + } + g.wheel.mu.Unlock() +} + +// MouseHWheel sends a horizontal scroll event (trackpad two-finger horizontal +// pan or tilt-wheel). delta is the rotation amount in physical notches; +// positive values scroll right, negative scroll left. Shares the vertical +// axis's coalescing window and timer. +func (g *RdpClient) MouseHWheel(delta float64) { + if !g.eventReady.Load() { + return + } + g.flushMouseMove() + + const wheelDelta = 120 + g.wheel.mu.Lock() + g.wheel.haccum += delta * wheelDelta + if g.wheel.haccum == 0 { + g.wheel.mu.Unlock() + return + } + now := time.Now() + since := now.Sub(g.wheel.lastTx) + if since >= mouseCoalesceInterval { + g.sendWheelLocked(now) + g.wheel.mu.Unlock() + return + } + if g.wheel.timer == nil { + delay := mouseCoalesceInterval - since + g.wheel.timer = time.AfterFunc(delay, g.flushWheelTimer) + } + g.wheel.mu.Unlock() +} + +// flushWheel sends any pending wheel event synchronously. Called before any +// non-wheel input event to preserve server-side ordering. +func (g *RdpClient) flushWheel() { + g.wheel.mu.Lock() + if g.wheel.timer != nil { + g.wheel.timer.Stop() + g.wheel.timer = nil + } + if g.wheel.accum != 0 { + g.sendWheelLocked(time.Now()) + } + g.wheel.mu.Unlock() +} + +// flushWheelTimer is the time.AfterFunc callback for wheel coalescing. +func (g *RdpClient) flushWheelTimer() { + g.wheel.mu.Lock() + g.wheel.timer = nil + if g.wheel.accum != 0 && g.eventReady.Load() { + g.sendWheelLocked(time.Now()) + } + g.wheel.mu.Unlock() +} + +// sendWheelLocked must be called with wheel.mu held. +// Modelled on FreeRDP's send_mouse_wheel in client/SDL/SDL2/sdl_touch.cpp. +func (g *RdpClient) sendWheelLocked(now time.Time) { + // Truncate the accumulated float to a whole WHEEL_DELTA integer; keep the + // fractional remainder so sub-notch trackpad movements aren't discarded. + iaccum := int(g.wheel.accum) + g.wheel.accum -= float64(iaccum) + haccum := int(g.wheel.haccum) + g.wheel.haccum -= float64(haccum) + g.wheel.lastTx = now + + // sendAxis emits 0xFF-capped wheel events for one axis. The WheelRotation + // field is 9 bits. Bits 0–7 hold the unsigned magnitude (max 0xFF per + // event); bit 8 is the sign (PTRFLAGS_WHEEL_NEGATIVE). For negative values + // the receiver computes -(0x100 - bits[0:7]), so we must store the 9-bit + // two's-complement form, not the raw magnitude (same loop as FreeRDP). + sendAxis := func(baseFlags uint16, v int) { + negative := v < 0 + if negative { + v = -v + } + if negative { + baseFlags |= uint16(pdu.PTRFLAGS_WHEEL_NEGATIVE) + } + for v > 0 { + cval := min(v, 0xFF) + v -= cval + if negative { + g.wheel.pdu.PointerFlags = (baseFlags & 0xFF00) | uint16(0x100-cval) + } else { + g.wheel.pdu.PointerFlags = baseFlags | uint16(cval) + } + g.pdu.SendInputEvents(pdu.INPUT_EVENT_MOUSE, g.wheel.pduBuf[:]) + } + } + + if iaccum != 0 { + sendAxis(uint16(pdu.PTRFLAGS_WHEEL), iaccum) + } + if haccum != 0 { + // 水平轴:正值 = 向右滚动(不设 NEGATIVE 位 = 远端向右)。 + sendAxis(uint16(pdu.PTRFLAGS_HWHEEL), haccum) + } + if iaccum != 0 || haccum != 0 { + g.notifyGfxLocalInput() + } +} + +func (g *RdpClient) MouseUp(button int, x, y int) { + if !g.eventReady.Load() { + return + } + slog.Debug("MouseUp", "x", x, "y", y, "button", button) + g.flushMouseMove() + g.flushWheel() + p := &pdu.PointerEvent{} + p.PointerFlags = mouseButtonFlag(button) + p.XPos = uint16(x) + p.YPos = uint16(y) + g.pdu.SendInputEvents(pdu.INPUT_EVENT_MOUSE, []pdu.InputEventsInterface{p}) + g.notifyGfxLocalInput() +} + +func (g *RdpClient) MouseDown(button int, x, y int) { + if !g.eventReady.Load() { + return + } + slog.Debug("MouseDown", "x", x, "y", y, "button", button) + g.flushMouseMove() + g.flushWheel() + p := &pdu.PointerEvent{} + p.PointerFlags = pdu.PTRFLAGS_DOWN | mouseButtonFlag(button) + p.XPos = uint16(x) + p.YPos = uint16(y) + g.pdu.SendInputEvents(pdu.INPUT_EVENT_MOUSE, []pdu.InputEventsInterface{p}) + g.notifyGfxLocalInput() +} + +// SetResolution requests a desktop resolution change via the MS-RDPEDISP +// Display Update Virtual Channel. The server will reshape the desktop to the +// given dimensions and send a fresh RDPGFX ResetGraphics command. +// +// width must be even and both width and height must be >= 200. +// This method is a no-op when the RDPEDISP channel has not been established +// (e.g. when the server does not support it). +func (g *RdpClient) SetResolution(width, height int) { + if g.dispHandler == nil { + slog.Warn("SetResolution: RDPEDISP channel not available") + return + } + w := uint32(width) + if w%2 != 0 { + w++ + } + w = max(w, 200) + h := uint32(height) + h = max(h, 200) + g.dispHandler.SendMonitorLayout([]rdpedisp.Monitor{ + { + Flags: rdpedisp.MonitorFlagPrimary, + Left: 0, + Top: 0, + Width: w, + Height: h, + PhysicalWidth: 0, + PhysicalHeight: 0, + Orientation: 0, + DesktopScaleFactor: 100, + DeviceScaleFactor: 100, + }, + }) + slog.Debug("SetResolution", "width", w, "height", h) +} + +// SetQueueDepthHint controls the frame-rate and encoding quality reported to +// the server via the RDPGFX FRAME_ACKNOWLEDGE queueDepth field +// (MS-RDPEGFX 2.2.2.8). +// +// A higher value signals a larger client decode backlog, causing the server to +// slow down or reduce H.264/RFX encoding quality. 0 (default) means "report +// the real decode-queue length" — no artificial throttling. +// DebugSurfacePixel 诊断:读第一个 mapped surface 上 (x,y) 的 BGRA +func (g *RdpClient) DebugSurfacePixel(x, y int) (uint8, uint8, uint8, uint8, bool) { + return g.gfxHandler.DebugSurfacePixel(x, y) +} + +// GfxDiagStats 返回图形管线实时诊断指标:累计解码帧数、解码队列深度、 +// 帧间隔(毫秒)与单消息解码耗时 EMA(微秒)。见 GfxHandler.DiagStats。 +func (g *RdpClient) GfxDiagStats() (frames, qdepth, fintvMs, decUs int64) { + return g.gfxHandler.DiagStats() +} + +// CodecStats exposes cumulative surface-bitmap bytes per codec id for +// bandwidth diagnostics. Returns nil before a graphics session exists. +func (g *RdpClient) CodecStats() map[uint16]int64 { + if g.gfxHandler != nil { + return g.gfxHandler.CodecStats() + } + return nil +} + +// Typical values: 0 = off, 10–50 = moderate throttle, 100+ = heavy throttle. +// Use 0xFFFFFFFF to pause new frames entirely (the stream resumes when hint is +// reduced or cleared). +func (g *RdpClient) SetQueueDepthHint(depth uint32) { + if g.gfxHandler != nil { + g.gfxHandler.SetQueueDepthHint(depth) + } +} + +// SetAudioMode 设置声音重定向模式(mstsc 远程音频播放): +// "local"(本机播放,默认)| "none"(不播放:协商后丢弃,服务器静音)| +// "remote"(远端播放:不注册音频通道,服务器本机扬声器出声)。 +// 必须在 Login 前调用。 +func (g *RdpClient) SetAudioMode(mode string) { + g.audioMode = mode +} + +// SetDriveRedirect 启用驱动器重定向(MS-RDPEFS):远端将看到一个只读的 +// 重定向设备,内容由 SetFilesystem 桥接的异步文件系统提供。必须在 Login +// 前调用。 +func (g *RdpClient) SetDriveRedirect(enabled bool) { + g.driveRedirect = enabled +} + +// SetDriveLabel 设置重定向卷的卷标(远端 Explorer 显示名)。必须在 Login +// 前调用;空值取 "local"。 +func (g *RdpClient) SetDriveLabel(label string) { + g.driveLabel = label +} + +// SetDriveFilesystem 挂接异步文件系统桥。处理器在 Login(doLogin)内才 +// 创建——这里先暂存字段,创建时补挂(同 SetGfxCacheStore 的时序教训)。 +func (g *RdpClient) SetDriveFilesystem(fs rdpdr.Filesystem) { + g.driveFS = fs +} + +// Rdpdr 返回驱动器重定向处理器(未启用时为 nil)。桥接层用它接收完成回调。 +func (g *RdpClient) Rdpdr() *rdpdr.Handler { return g.rdpdrHandler } + +// SetSessionColorDepth 请求会话颜色位数(16/24/32,其它值按 32 处理)。 +// 仅影响传统位图管线;RDPGFX 会话表面恒为 32bpp。必须在 Login 前调用。 +func (g *RdpClient) SetSessionColorDepth(bpp int) { + g.colorDepth = bpp +} + +// SetPerformanceFlags 覆盖 Client Info PDU 的 performanceFlags(体验选项)。 +// 禁用类位置位 = 关闭(PERF_DISABLE_WALLPAPER 等),0 = 视觉全开默认。 +// 必须在 Login 前调用。 +func (g *RdpClient) SetPerformanceFlags(flags uint32) { + g.perfFlags = flags + g.perfFlagsSet = true +} + +// RequestKeyframe asks the server to send a fresh full-screen IDR keyframe via +// the SuppressOutput off→on toggle (SendForceRefresh). It is the same request +// the RDPGFX decoder issues internally when it stalls, exposed publicly so the +// frontend's black-screen watchdog can recover an initial all-black session — +// where the server sent its first IDR before the desktop finished painting +// (decoded as a black warm-up frame and dropped) and then went idle — without +// the cost of a full reconnect. Safe to call from the render loop goroutine. +func (g *RdpClient) RequestKeyframe() { + if g.closed.Load() { + return + } + if g.pdu != nil { + g.pdu.SendForceRefresh(uint16(g.width), uint16(g.height)) + } +} + +func (g *RdpClient) Reconnect(width, height int) error { + if g.closed.Load() { + return fmt.Errorf("client is closed") + } + + g.reconnectMu.Lock() + defer g.reconnectMu.Unlock() + + g.reconnecting.Store(true) + defer func() { g.reconnecting.Store(false) }() + + slog.Debug("Reconnect", "width", width, "height", height) + g.closeTransport() + g.width = width + g.height = height + g.eventReady.Store(false) + + const maxRetries = 3 + for attempt := 1; attempt <= maxRetries; attempt++ { + // No delay on the first attempt: the transport was already closed above + // so the server has already started session teardown. Use exponential + // backoff (1s, 2s) only for retries after a failed login. + delay := time.Duration(0) + if attempt > 1 { + delay = time.Duration(1< 30 { + return CHANNEL_RC_TOO_MANY_CHANNELS + } + + if channels.connected { + return CHANNEL_RC_ALREADY_CONNECTED + } + + for i := range pChannel { + pChannelDef := &pChannel[i] + if getChannelOpenDataByName(channels, pChannelDef.Name) == nil { + return CHANNEL_RC_BAD_CHANNEL + } + } + + pChannelClientData = &channels.clientDataList[channels.clientDataCount] + pChannelClientData.pChannelInitEventProcEx = pChannelInitEventProcEx + pChannelClientData.pInitHandle = pInitHandle + pChannelClientData.lpUserParam = lpUserParam + channels.clientDataCount++ + + pChannelInitData.pInterface = clientContext + + for i := range pChannel { + pChannelDef := &pChannel[i] + var pChannelOpenData ChannelOpenData + + pChannelOpenData.OpenHandle = atomic.AddUint32(&openHandleSeq, 1) + pChannelOpenData.channels = channels + pChannelOpenData.lpUserParam = lpUserParam + if _, ok := pChannelInitData.openDataMap[pChannelOpenData.OpenHandle]; ok { + return CHANNEL_RC_INITIALIZATION_ERROR + } + + pChannelOpenData.flags = 1 + pChannelOpenData.name = pChannelDef.Name + pChannelOpenData.options = pChannelDef.Options + pChannelInitData.openDataMap[pChannelOpenData.OpenHandle] = &pChannelOpenData + channels.openDataList = append(channels.openDataList, pChannelOpenData) + channels.openDataCount++ + } + + return CHANNEL_RC_OK +} +func getChannelOpenDataByName(channel *rdpChannels, name string) *ChannelOpenData { + for _, v := range channel.openDataList { + if strings.EqualFold(name, v.name) { + return &v + } + } + return nil +} +func RdpVirtualChannelOpenEx(pInitHandle interface{}, pOpenHandle *uint32, pChannelName string, + pChannelOpenEventProcEx *CHANNEL_OPEN_EVENT_EX_FN) uint { + pChannelInitData := pInitHandle.(*ChannelInitData) + channels := pChannelInitData.channels + pInterface := pChannelInitData.pInterface + + if pOpenHandle == nil { + return CHANNEL_RC_BAD_CHANNEL_HANDLE + } + if pChannelOpenEventProcEx == nil { + return CHANNEL_RC_BAD_PROC + } + + if !channels.connected { + return CHANNEL_RC_NOT_CONNECTED + } + + pChannelOpenData := getChannelOpenDataByName(channels, pChannelName) + + if pChannelOpenData == nil { + return CHANNEL_RC_UNKNOWN_CHANNEL_NAME + } + + if pChannelOpenData.flags == 2 { + return CHANNEL_RC_ALREADY_OPEN + } + + pChannelOpenData.flags = 2 /* open */ + pChannelOpenData.pInterface = pInterface + pChannelOpenData.pChannelOpenEventProcEx = pChannelOpenEventProcEx + *pOpenHandle = pChannelOpenData.OpenHandle + return CHANNEL_RC_OK +} +func RdpVirtualChannelCloseEx(pInitHandle interface{}, openHandle uint32) uint { + if pInitHandle == nil { + return CHANNEL_RC_BAD_INIT_HANDLE + } + pChannelInitData := pInitHandle.(*ChannelInitData) + pChannelOpenData := pChannelInitData.openDataMap[openHandle] + + if pChannelOpenData == nil { + return CHANNEL_RC_BAD_CHANNEL_HANDLE + } + + if pChannelOpenData.flags != 2 { + return CHANNEL_RC_NOT_OPEN + } + + pChannelOpenData.flags = 0 + + return CHANNEL_RC_OK +} +func RdpVirtualChannelWriteEx(pInitHandle interface{}, openHandle uint32, + pData interface{}, dataLength uint32, + pUserData interface{}) uint { + + //wMessage message; + + if pInitHandle == nil { + return CHANNEL_RC_BAD_INIT_HANDLE + } + + pChannelInitData := pInitHandle.(*ChannelInitData) + channels := pChannelInitData.channels + + if channels == nil { + return CHANNEL_RC_BAD_CHANNEL_HANDLE + } + + pChannelOpenData := pChannelInitData.openDataMap[openHandle] + if pChannelOpenData == nil { + return CHANNEL_RC_BAD_CHANNEL_HANDLE + } + + if !channels.connected { + return CHANNEL_RC_NOT_CONNECTED + } + + if pData == nil { + return CHANNEL_RC_NULL_DATA + } + + if dataLength == 0 { + return CHANNEL_RC_ZERO_LENGTH + } + + if pChannelOpenData.flags != 2 { + return CHANNEL_RC_NOT_OPEN + } + + pChannelOpenEvent := new(ChannelOpenEvent) + + if pChannelOpenEvent == nil { + return CHANNEL_RC_NO_MEMORY + + } + + pChannelOpenEvent.Data = pData + pChannelOpenEvent.DataLength = dataLength + pChannelOpenEvent.UserData = pUserData + pChannelOpenEvent.pChannelOpenData = pChannelOpenData + /*message.context = channels; + message.id = 0; + message.wParam = pChannelOpenEvent; + message.lParam = NULL; + message.Free = channel_queue_message_free; + + if (!MessageQueue_Dispatch(channels->queue, &message)) + { + free(pChannelOpenEvent); + return CHANNEL_RC_NO_MEMORY; + }*/ + + return CHANNEL_RC_OK +} diff --git a/plugin/channel.go b/plugin/channel.go new file mode 100644 index 0000000..9a2ae0b --- /dev/null +++ b/plugin/channel.go @@ -0,0 +1,308 @@ +package plugin + +import ( + "bytes" + "encoding/binary" + "fmt" + "log/slog" + "slices" + "sync" + "unsafe" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" +) + +// chunkBufPool reuses the header+payload buffers for outbound channel chunks. +// Each buffer is pre-sized to the largest possible chunk (8-byte header + 1600 bytes data). +var chunkBufPool = sync.Pool{ + New: func() any { return make([]byte, 0, 8+CHANNEL_CHUNK_LENGTH) }, +} + +const ( + CHANNEL_RC_OK = 0 + CHANNEL_RC_ALREADY_INITIALIZED = 1 + CHANNEL_RC_NOT_INITIALIZED = 2 + CHANNEL_RC_ALREADY_CONNECTED = 3 + CHANNEL_RC_NOT_CONNECTED = 4 + CHANNEL_RC_TOO_MANY_CHANNELS = 5 + CHANNEL_RC_BAD_CHANNEL = 6 + CHANNEL_RC_BAD_CHANNEL_HANDLE = 7 + CHANNEL_RC_NO_BUFFER = 8 + CHANNEL_RC_BAD_INIT_HANDLE = 9 + CHANNEL_RC_NOT_OPEN = 10 + CHANNEL_RC_BAD_PROC = 11 + CHANNEL_RC_NO_MEMORY = 12 + CHANNEL_RC_UNKNOWN_CHANNEL_NAME = 13 + CHANNEL_RC_ALREADY_OPEN = 14 + CHANNEL_RC_NOT_IN_VIRTUALCHANNELENTRY = 15 + CHANNEL_RC_NULL_DATA = 16 + CHANNEL_RC_ZERO_LENGTH = 17 + CHANNEL_RC_INVALID_INSTANCE = 18 + CHANNEL_RC_UNSUPPORTED_VERSION = 19 + CHANNEL_RC_INITIALIZATION_ERROR = 20 +) +const ( + VIRTUAL_CHANNEL_VERSION_WIN2000 = 1 +) + +const ( + CHANNEL_EVENT_INITIALIZED = 0 + CHANNEL_EVENT_CONNECTED = 1 + CHANNEL_EVENT_V1_CONNECTED = 2 + CHANNEL_EVENT_DISCONNECTED = 3 + CHANNEL_EVENT_TERMINATED = 4 + CHANNEL_EVENT_REMOTE_CONTROL_START = 5 + CHANNEL_EVENT_REMOTE_CONTROL_STOP = 6 + CHANNEL_EVENT_ATTACHED = 7 + CHANNEL_EVENT_DETACHED = 8 + CHANNEL_EVENT_DATA_RECEIVED = 10 + CHANNEL_EVENT_WRITE_COMPLETE = 11 + CHANNEL_EVENT_WRITE_CANCELLED = 12 +) + +const ( + CHANNEL_OPTION_INITIALIZED = 0x80000000 + CHANNEL_OPTION_ENCRYPT_RDP = 0x40000000 + CHANNEL_OPTION_ENCRYPT_SC = 0x20000000 + CHANNEL_OPTION_ENCRYPT_CS = 0x10000000 + CHANNEL_OPTION_PRI_HIGH = 0x08000000 + CHANNEL_OPTION_PRI_MED = 0x04000000 + CHANNEL_OPTION_PRI_LOW = 0x02000000 + CHANNEL_OPTION_COMPRESS_RDP = 0x00800000 + CHANNEL_OPTION_COMPRESS = 0x00400000 + CHANNEL_OPTION_SHOW_PROTOCOL = 0x00200000 + CHANNEL_OPTION_REMOTE_CONTROL_PERSISTENT = 0x00100000 +) + +type ChannelDef struct { + Name string + Options uint32 +} +type CHANNEL_INIT_EVENT_EX_FN func(lpUserParam any, + pInitHandle any, event uint, pData uintptr, dataLength uint) +type VIRTUALCHANNELINITEX func(lpUserParam any, clientContext any, + pInitHandle any, pChannel []ChannelDef, + channelCount int, versionRequested uint32, + pChannelInitEventProcEx CHANNEL_INIT_EVENT_EX_FN) uint + +type CHANNEL_OPEN_EVENT_EX_FN func(lpUserParam uintptr, + openHandle uint32, event uint, + pData uintptr, dataLength uint32, totalLength uint32, dataFlags uint32) +type VIRTUALCHANNELOPENEX func(pInitHandle any, pOpenHandle *uint32, + pChannelName string, + pChannelOpenEventProcEx *CHANNEL_OPEN_EVENT_EX_FN) uint + +type VIRTUALCHANNELCLOSEEX func(pInitHandle any, openHandle uint32) uint + +type VIRTUALCHANNELWRITEEX func(pInitHandle any, openHandle uint32, pData any, + dataLength uint32, pUserData any) uint + +type ChannelEntryPointsEx struct { + CbSize uint32 + ProtocolVersion uint32 + PVirtualChannelInitEx VIRTUALCHANNELINITEX + PVirtualChannelOpenEx VIRTUALCHANNELOPENEX + PVirtualChannelCloseEx VIRTUALCHANNELCLOSEEX + PVirtualChannelWriteEx VIRTUALCHANNELWRITEEX +} + +func NewChannelEntryPointsEx() *ChannelEntryPointsEx { + e := &ChannelEntryPointsEx{} + e.CbSize = uint32(unsafe.Sizeof(e)) + e.ProtocolVersion = VIRTUAL_CHANNEL_VERSION_WIN2000 + return e +} + +type VIRTUALCHANNELENTRYEX func(pEntryPointsEx *ChannelEntryPointsEx, + pInitHandle any) error + +/* +type ChannelEntryPoints struct { + CbSize uint32 + ProtocolVersion uint32 + PVirtualChannelInit PVIRTUALCHANNELINIT + PVirtualChannelOpen PVIRTUALCHANNELOPEN + PVirtualChannelClose PVIRTUALCHANNELCLOSE + PVirtualChannelWrite PVIRTUALCHANNELWRITE +} +typedef VOID VCAPITYPE CHANNEL_INIT_EVENT_FN(LPVOID pInitHandle, + UINT event, LPVOID pData, UINT dataLength); + +typedef CHANNEL_INIT_EVENT_FN* PCHANNEL_INIT_EVENT_FN; +typedef VOID VCAPITYPE CHANNEL_OPEN_EVENT_FN(DWORD openHandle, UINT event, + LPVOID pData, UINT32 dataLength, UINT32 totalLength, UINT32 dataFlags); + +typedef CHANNEL_OPEN_EVENT_FN* PCHANNEL_OPEN_EVENT_FN; +typedef UINT VCAPITYPE VIRTUALCHANNELINIT(LPVOID* ppInitHandle, PCHANNEL_DEF pChannel, + INT channelCount, ULONG versionRequested, + PCHANNEL_INIT_EVENT_FN pChannelInitEventProc); +typedef VIRTUALCHANNELINIT* PVIRTUALCHANNELINIT; + +typedef UINT VCAPITYPE VIRTUALCHANNELOPEN(LPVOID pInitHandle, LPDWORD pOpenHandle, + PCHAR pChannelName, + PCHANNEL_OPEN_EVENT_FN pChannelOpenEventProc); + +typedef VIRTUALCHANNELOPEN* PVIRTUALCHANNELOPEN; + +typedef UINT VCAPITYPE VIRTUALCHANNELCLOSE(DWORD openHandle); +typedef VIRTUALCHANNELCLOSE* PVIRTUALCHANNELCLOSE; + +typedef UINT VCAPITYPE VIRTUALCHANNELWRITE(DWORD openHandle, LPVOID pData, ULONG dataLength, + LPVOID pUserData); +typedef VIRTUALCHANNELWRITE* PVIRTUALCHANNELWRITE; + +typedef UINT VCAPITYPE VIRTUALCHANNELINITEX(LPVOID lpUserParam, LPVOID clientContext, + LPVOID pInitHandle, PCHANNEL_DEF pChannel, + INT channelCount, ULONG versionRequested, + PCHANNEL_INIT_EVENT_EX_FN pChannelInitEventProcEx); +typedef VIRTUALCHANNELINITEX* PVIRTUALCHANNELINITEX; + +typedef UINT VCAPITYPE VIRTUALCHANNELOPENEX(LPVOID pInitHandle, LPDWORD pOpenHandle, + PCHAR pChannelName, + PCHANNEL_OPEN_EVENT_EX_FN pChannelOpenEventProcEx); +typedef VIRTUALCHANNELOPENEX* PVIRTUALCHANNELOPENEX; + + +typedef UINT VCAPITYPE VIRTUALCHANNELCLOSEEX(LPVOID pInitHandle, DWORD openHandle); +typedef VIRTUALCHANNELCLOSEEX* PVIRTUALCHANNELCLOSEEX; + +typedef UINT VCAPITYPE VIRTUALCHANNELWRITEEX(LPVOID pInitHandle, DWORD openHandle, LPVOID pData, + ULONG dataLength, LPVOID pUserData); +typedef VIRTUALCHANNELWRITEEX* PVIRTUALCHANNELWRITEEX; +*/ + +// static channel name +const ( + CLIPRDR_SVC_CHANNEL_NAME = "cliprdr" // clipboard + RDPDR_SVC_CHANNEL_NAME = "rdpdr" // device redirection + RDPSND_SVC_CHANNEL_NAME = "rdpsnd" // sound + RAIL_SVC_CHANNEL_NAME = "rail" // remote appication + DRDYNVC_SVC_CHANNEL_NAME = "drdynvc" // dynamic virtual channel + REMDESK_SVC_CHANNEL_NAME = "remdesk" // remote assistance +) + +const ( + RDPGFX_DVC_CHANNEL_NAME = "Microsoft::Windows::RDS::Graphics" // Graphics Extension +) + +var StaticVirtualChannels = map[string]int{ + CLIPRDR_SVC_CHANNEL_NAME: CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | + CHANNEL_OPTION_COMPRESS_RDP | CHANNEL_OPTION_SHOW_PROTOCOL, + RDPDR_SVC_CHANNEL_NAME: CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | CHANNEL_OPTION_COMPRESS_RDP, + RDPSND_SVC_CHANNEL_NAME: CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | + CHANNEL_OPTION_COMPRESS_RDP | CHANNEL_OPTION_SHOW_PROTOCOL, + RAIL_SVC_CHANNEL_NAME: CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | + CHANNEL_OPTION_COMPRESS_RDP | CHANNEL_OPTION_SHOW_PROTOCOL, +} + +const ( + CHANNEL_CHUNK_LENGTH = 1600 + CHANNEL_FLAG_FIRST = 0x01 + CHANNEL_FLAG_LAST = 0x02 + CHANNEL_FLAG_SHOW_PROTOCOL = 0x10 +) + +type ChannelTransport interface { + GetType() (string, uint32) + Sender(core.ChannelSender) + Process(s []byte) +} +type ChannelClient struct { + ChannelDef + t ChannelTransport +} + +type Channels struct { + emission.Emitter + channels map[string]ChannelClient + transport core.Transport + buff *bytes.Buffer + channelSender core.ChannelSender +} + +func NewChannels(t core.Transport) *Channels { + c := &Channels{ + Emitter: *emission.NewEmitter(), + channels: make(map[string]ChannelClient, 20), + transport: t, + buff: &bytes.Buffer{}, + } + t.On("channel", c.process) + return c +} + +func (c *Channels) SetChannelSender(f core.ChannelSender) { + c.channelSender = f +} +func (c *Channels) Register(t ChannelTransport) { + name, option := t.GetType() + _, ok := c.channels[name] + if ok { + slog.Warn("Already register", "channel", name) + return + } + t.Sender(c) + c.channels[name] = ChannelClient{ChannelDef{name, option}, t} +} + +func (c *Channels) SendToChannel(channel string, s []byte) (int, error) { + cli, ok := c.channels[channel] + if !ok { + slog.Warn("No register", "channel", channel) + return 0, fmt.Errorf("No register channel: %s", channel) + } + totalLen := len(s) + baseFlag := uint32(0) + if cli.Options&CHANNEL_OPTION_SHOW_PROTOCOL != 0 { + baseFlag |= CHANNEL_FLAG_SHOW_PROTOCOL + } + buf := chunkBufPool.Get().([]byte) + first := true + remaining := totalLen + for chunk := range slices.Chunk(s, CHANNEL_CHUNK_LENGTH) { + flag := baseFlag + if first { + flag |= CHANNEL_FLAG_FIRST + first = false + } + remaining -= len(chunk) + if remaining == 0 { + flag |= CHANNEL_FLAG_LAST + } + slog.Debug("SendToChannel", "len", len(chunk), "flag", flag) + buf = buf[:8+len(chunk)] + binary.LittleEndian.PutUint32(buf[0:], uint32(totalLen)) + binary.LittleEndian.PutUint32(buf[4:], flag) + copy(buf[8:], chunk) + c.channelSender.SendToChannel(channel, buf) + } + chunkBufPool.Put(buf[:0]) + return 0, nil +} + +func (c *Channels) process(channel string, s []byte) { + cli, ok := c.channels[channel] + if !ok { + slog.Warn("process No found channel", "channel", channel) + return + } + if len(s) < 8 { + return + } + // Parse 8-byte header directly: totalLen(4) + flags(4) + flags := uint32(s[4]) | uint32(s[5])<<8 | uint32(s[6])<<16 | uint32(s[7])<<24 + payload := s[8:] + if flags&CHANNEL_FLAG_FIRST == 0 || flags&CHANNEL_FLAG_LAST == 0 { + if flags&CHANNEL_FLAG_FIRST != 0 { + c.buff.Reset() + } + c.buff.Write(payload) + if flags&CHANNEL_FLAG_LAST == 0 { + return + } + cli.t.Process(c.buff.Bytes()) + } else { + cli.t.Process(payload) + } +} diff --git a/plugin/cliprdr/cliprdr.go b/plugin/cliprdr/cliprdr.go new file mode 100644 index 0000000..8119dda --- /dev/null +++ b/plugin/cliprdr/cliprdr.go @@ -0,0 +1,265 @@ +package cliprdr + +import ( + "bytes" + "fmt" + "log/slog" + "unicode/utf16" + + "github.com/lunixbochs/struc" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +type CliprdrClient struct { + w core.ChannelSender + useLongFormatNames bool + streamFileClipEnabled bool + fileClipNoFilePaths bool + canLockClipData bool + hasHugeFileSupport bool + formatIdMap map[uint32]uint32 + reply chan []byte +} + +func NewCliprdrClient() *CliprdrClient { + c := &CliprdrClient{ + formatIdMap: make(map[uint32]uint32, 20), + reply: make(chan []byte, 100), + } + + go ClipWatcher(c) + + return c +} + +func (c *CliprdrClient) Sender(f core.ChannelSender) { + c.w = f +} +func (c *CliprdrClient) GetType() (string, uint32) { + return ChannelName, ChannelOption +} + +func (c *CliprdrClient) Process(s []byte) { + r := bytes.NewReader(s) + + msgType, _ := core.ReadUint16LE(r) + flag, _ := core.ReadUint16LE(r) + length, _ := core.ReadUInt32LE(r) + slog.Debug(fmt.Sprintf("cliprdr: type=0x%x flag=%d length=%d, all=%d", msgType, flag, length, r.Len())) + + b, _ := core.ReadBytes(int(length), r) + + switch msgType { + case CB_CLIP_CAPS: + slog.Debug("CB_CLIP_CAPS") + c.processClipCaps(b) + + case CB_MONITOR_READY: + slog.Debug("CB_MONITOR_READY") + c.processMonitorReady(b) + + case CB_FORMAT_LIST: + slog.Debug("CB_FORMAT_LIST") + c.processFormatList(b) + + case CB_FORMAT_LIST_RESPONSE: + slog.Debug("CB_FORMAT_LIST_RESPONSE") + c.processFormatListResponse(flag, b) + + case CB_FORMAT_DATA_REQUEST: + slog.Debug("CB_FORMAT_DATA_REQUEST") + c.processFormatDataRequest(b) + + case CB_FORMAT_DATA_RESPONSE: + slog.Debug("CB_FORMAT_DATA_RESPONSE") + c.processFormatDataResponse(flag, b) + + case CB_FILECONTENTS_REQUEST: + slog.Debug("CB_FILECONTENTS_REQUEST") + c.processFileContentsRequest(b) + + case CB_FILECONTENTS_RESPONSE: + slog.Debug("CB_FILECONTENTS_RESPONSE") + c.processFileContentsResponse(flag, b) + + case CB_LOCK_CLIPDATA: + slog.Debug("CB_LOCK_CLIPDATA") + c.processLockClipData(b) + + case CB_UNLOCK_CLIPDATA: + slog.Debug("CB_UNLOCK_CLIPDATA") + c.processUnlockClipData(b) + + default: + slog.Error(fmt.Sprintf("type 0x%x not supported", msgType)) + } +} +func (c *CliprdrClient) processClipCaps(b []byte) { + r := bytes.NewReader(b) + var cp CliprdrCapabilitiesPDU + err := struc.Unpack(r, &cp) + if err != nil { + slog.Error("Failed to unpack", "error", err) + return + } + slog.Debug(fmt.Sprintf("Capabilities:%+v", cp)) + c.useLongFormatNames = cp.CapabilitySets[0].GeneralFlags&CB_USE_LONG_FORMAT_NAMES != 0 + c.streamFileClipEnabled = cp.CapabilitySets[0].GeneralFlags&CB_STREAM_FILECLIP_ENABLED != 0 + c.fileClipNoFilePaths = cp.CapabilitySets[0].GeneralFlags&CB_FILECLIP_NO_FILE_PATHS != 0 + c.canLockClipData = cp.CapabilitySets[0].GeneralFlags&CB_CAN_LOCK_CLIPDATA != 0 + c.hasHugeFileSupport = cp.CapabilitySets[0].GeneralFlags&CB_HUGE_FILE_SUPPORT_ENABLED != 0 + slog.Debug("UseLongFormatNames", "value", c.useLongFormatNames) + slog.Debug("StreamFileClipEnabled", "value", c.streamFileClipEnabled) + slog.Debug("FileClipNoFilePaths", "value", c.fileClipNoFilePaths) + slog.Debug("CanLockClipData", "value", c.canLockClipData) + slog.Debug("HasHugeFileSupport", "value", c.hasHugeFileSupport) +} + +func (c *CliprdrClient) processMonitorReady(b []byte) { + c.sendClientCapabilitiesPDU() + c.sendFormatListPDU() +} + +func (c *CliprdrClient) processFormatList(b []byte) { + EmptyClipboard() + fl, _ := c.readFormatList(b) + slog.Debug("numFormats", "count", fl.NumFormats) + c.sendFormatListResponse(CB_RESPONSE_OK) +} + +func (c *CliprdrClient) processFormatListResponse(flag uint16, b []byte) { + if flag != CB_RESPONSE_OK { + slog.Error("Format List Response Failed") + return + } + slog.Debug("Format List Response OK") +} + +func (c *CliprdrClient) processFormatDataRequest(b []byte) { + r := bytes.NewReader(b) + _, _ = core.ReadUInt32LE(r) // requestId + + buff := &bytes.Buffer{} + // Text-only: directly get clipboard data for any text format + data := GetClipboardText() + slog.Debug("clipboard data", "content", data) + buff.Write(core.UnicodeEncode(data)) + buff.Write([]byte{0, 0}) + + c.sendFormatDataResponse(buff.Bytes()) +} +func (c *CliprdrClient) processFormatDataResponse(flag uint16, b []byte) { + if flag != CB_RESPONSE_OK { + slog.Error("Format Data Response Failed") + } + c.reply <- b +} + +func (c *CliprdrClient) processFileContentsRequest(b []byte) { + // Text-only mode doesn't support file transfer + slog.Debug("File transfer not supported in text-only mode") +} + +func (c *CliprdrClient) processFileContentsResponse(flag uint16, b []byte) { + // Text-only mode doesn't support file transfer +} +func (c *CliprdrClient) processLockClipData(b []byte) { + r := bytes.NewReader(b) + var l CliprdrCtrlClipboardData + l.ClipDataId, _ = core.ReadUInt32LE(r) +} +func (c *CliprdrClient) processUnlockClipData(b []byte) { + r := bytes.NewReader(b) + var l CliprdrCtrlClipboardData + l.ClipDataId, _ = core.ReadUInt32LE(r) + +} + +func (c *CliprdrClient) sendClientCapabilitiesPDU() { + slog.Debug("Send Client Clipboard Capabilities PDU (text-only mode)") + var cs CliprdrGeneralCapabilitySet + cs.CapabilitySetLength = 12 + cs.CapabilitySetType = CB_CAPSTYPE_GENERAL + cs.Version = CB_CAPS_VERSION_2 + // Text-only mode: only use long format names + cs.GeneralFlags = CB_USE_LONG_FORMAT_NAMES + body := &bytes.Buffer{} + core.WriteUInt16LE(1, body) // cCapabilitiesSets + core.WriteUInt16LE(0, body) // pad + struc.Pack(body, cs) + sendClipPDU(c.w, CB_CLIP_CAPS, 0, body.Bytes()) +} + +func (c *CliprdrClient) sendTemporaryDirectoryPDU() { + slog.Debug("Send Temporary Directory PDU (ignored in text-only mode)") +} + +func (c *CliprdrClient) sendFormatListPDU() { + slog.Debug("Send Format List PDU (text formats only)") + formats := GetFormatList() + slog.Debug("available formats", "count", len(formats), "formats", formats) + + body := &bytes.Buffer{} + for _, v := range formats { + core.WriteUInt32LE(v.FormatId, body) + if v.FormatName == "" { + core.WriteUInt16LE(0, body) + } else { + n := core.UnicodeEncode(v.FormatName) + core.WriteBytes(n, body) + body.Write([]byte{0, 0}) + } + } + sendClipPDU(c.w, CB_FORMAT_LIST, 0, body.Bytes()) +} + +func (c *CliprdrClient) readFormatList(b []byte) (*CliprdrFormatList, bool) { + r := bytes.NewReader(b) + fs := make([]CliprdrFormat, 0, 20) + var numFormats uint32 = 0 + c.formatIdMap = make(map[uint32]uint32, 0) + for r.Len() > 0 { + formatId, _ := core.ReadUInt32LE(r) + bs := make([]uint16, 0, 20) + ln := r.Len() + for range ln { + b, _ := core.ReadUint16LE(r) + if b == 0 { + break + } + bs = append(bs, b) + } + name := string(utf16.Decode(bs)) + slog.Debug(fmt.Sprintf("Format:%d Name:<%s>", formatId, name)) + if name != "" { + localId := RegisterClipboardFormat(name) + slog.Debug("format mapping", "local", localId, "remote", formatId) + c.formatIdMap[localId] = formatId + } else { + c.formatIdMap[formatId] = formatId + } + + numFormats++ + fs = append(fs, CliprdrFormat{formatId, name}) + } + + return &CliprdrFormatList{numFormats, fs}, false +} + +func (c *CliprdrClient) sendFormatListResponse(flags uint16) { + slog.Debug("Send Format List Response") + sendClipPDU(c.w, CB_FORMAT_LIST_RESPONSE, flags, nil) +} + +func (c *CliprdrClient) sendFormatDataRequest(id uint32) { + slog.Debug("Send Format Data Request") + body := &bytes.Buffer{} + core.WriteUInt32LE(id, body) + sendClipPDU(c.w, CB_FORMAT_DATA_REQUEST, 0, body.Bytes()) +} + +func (c *CliprdrClient) sendFormatDataResponse(b []byte) { + slog.Debug("Send Format Data Response") + sendClipPDU(c.w, CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, b) +} diff --git a/plugin/cliprdr/cliprdr_generic.go b/plugin/cliprdr/cliprdr_generic.go new file mode 100644 index 0000000..0b0a228 --- /dev/null +++ b/plugin/cliprdr/cliprdr_generic.go @@ -0,0 +1,181 @@ +// Generic, OS-independent clipboard support with text format only +package cliprdr + +import ( + "log/slog" +) + +// SimpleTextClipboard provides OS-independent text-only clipboard support +type SimpleTextClipboard struct { + textContent string +} + +var textClipboard = &SimpleTextClipboard{} + +// GetClipboardText returns the current text content +func GetClipboardText() string { + return textClipboard.textContent +} + +// SetClipboardText sets the clipboard text content +func SetClipboardText(text string) { + textClipboard.textContent = text + slog.Debug("Text clipboard set", "length", len(text)) +} + +// GetFormatList returns available formats (text only in generic mode) +func GetFormatList() []CliprdrFormat { + formatId := uint32(CF_UNICODETEXT) + formats := make([]CliprdrFormat, 0, 1) + formats = append(formats, CliprdrFormat{ + FormatId: formatId, + FormatName: "CF_UNICODETEXT", + }) + return formats +} + +// ClipWatcher is a no-op in generic mode +func ClipWatcher(c *CliprdrClient) { + slog.Debug("Generic clipboard watcher (text-only mode) started") + // In generic mode, we don't actively monitor system clipboard + // Format list is provided statically + select {} // Block indefinitely +} + +// Stub functions for compatibility (no-op in generic mode) + +func OpenClipboard(hwnd uintptr) bool { + return true +} + +func CloseClipboard() bool { + return true +} + +func CountClipboardFormats() int32 { + return 1 +} + +func IsClipboardFormatAvailable(id uint32) bool { + return id == CF_UNICODETEXT || id == CF_TEXT +} + +func EnumClipboardFormats(formatId uint32) uint32 { + if formatId == 0 { + return CF_UNICODETEXT + } + return 0 +} + +func GetClipboardFormatName(id uint32) string { + switch id { + case CF_TEXT: + return "CF_TEXT" + case CF_UNICODETEXT: + return "CF_UNICODETEXT" + default: + return "" + } +} + +func EmptyClipboard() bool { + textClipboard.textContent = "" + return true +} + +func RegisterClipboardFormat(format string) uint32 { + // Simple hash-based format ID generation + sum := uint32(0) + for _, c := range format { + sum = sum*31 + uint32(c) + } + return sum | 0xC000 // Set bit 14 for custom formats +} + +func IsClipboardOwner(h uintptr) bool { + return false // Not applicable in generic mode +} + +func HmemAlloc(data []byte) uintptr { + // In generic mode, return a pointer to the data + if len(data) == 0 { + return 0 + } + return uintptr(len(data)) +} + +func SetClipboardData(formatId uint32, hmem uintptr) bool { + // No-op in generic mode + return true +} + +func GetClipboardData(formatId uint32) string { + if formatId == CF_UNICODETEXT || formatId == CF_TEXT { + return GetClipboardText() + } + return "" +} + +func GlobalSize(hMem uintptr) uintptr { + return hMem +} + +func GlobalLock(hMem uintptr) uintptr { + return hMem +} + +func GlobalUnlock(hMem uintptr) { + // No-op +} + +func OleGetClipboard() *IDataObject { + return nil +} + +func OleSetClipboard(dataObject *IDataObject) bool { + return true +} + +func OleIsCurrentClipboard(dataObject *IDataObject) bool { + return true +} + +// GetFileNames and GetFileInfo are not supported in text-only generic mode. +func GetFileNames() []string { + return []string{} +} + +func GetFileInfo(sys any) (uint32, []byte, uint32, uint32) { + return 0, []byte{}, 0, 0 +} + +// IDataObject stub - not used in text-only mode +type IDataObject struct { + ptr uintptr +} + +type IUnknown struct { + ptr uintptr +} + +type FORMATETC struct { + CFormat uint32 + DvTargetDevice uintptr + Aspect uint32 + Index int32 + Tymed uint32 +} + +type STGMEDIUM struct { + Tymed uint32 + UnionMember uintptr + PUnkForRelease *IUnknown +} + +func (s *STGMEDIUM) Bytes() ([]byte, error) { + return []byte{}, nil +} + +func CreateDataObject(c *CliprdrClient) *IDataObject { + return nil +} diff --git a/plugin/cliprdr/cliprdr_image_test.go b/plugin/cliprdr/cliprdr_image_test.go new file mode 100644 index 0000000..3f4a3cb --- /dev/null +++ b/plugin/cliprdr/cliprdr_image_test.go @@ -0,0 +1,198 @@ +// cliprdr_image_test.go — 剪贴板图片链路的协议层单元测试(native 可跑) +package cliprdr_test + +import ( + "bytes" + "encoding/binary" + "image/png" + "sync" + "testing" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/plugin/cliprdr" +) + +// captureSender 捕获 handler 发出的 PDU,供断言 +type captureSender struct { + mu sync.Mutex + pdus [][]byte +} + +func (s *captureSender) GetType() (string, uint32) { return "cliprdr", 0 } +func (s *captureSender) Sender(f core.ChannelSender) {} +func (s *captureSender) Process(data []byte) {} + +func (s *captureSender) SendToChannel(name string, data []byte) (int, error) { + s.mu.Lock() + defer s.mu.Unlock() + pdu := make([]byte, len(data)) + copy(pdu, data) + s.pdus = append(s.pdus, pdu) + return len(data), nil +} + +func (s *captureSender) findByType(msgType uint16) [][]byte { + s.mu.Lock() + defer s.mu.Unlock() + var out [][]byte + for _, p := range s.pdus { + if len(p) >= 8 && binary.LittleEndian.Uint16(p[0:]) == msgType { + out = append(out, p) + } + } + return out +} + +// fakeDIB 构造一张 2x2 的 32bpp BI_RGB 底行优先 DIB +func fakeDIB() []byte { + const w, h = 2, 2 + dib := make([]byte, 40+w*4*h) + binary.LittleEndian.PutUint32(dib[0:4], 40) // biSize + binary.LittleEndian.PutUint32(dib[4:8], w) // width + binary.LittleEndian.PutUint32(dib[8:12], h) // height(正 = 底行优先) + binary.LittleEndian.PutUint16(dib[12:14], 1) // planes + binary.LittleEndian.PutUint16(dib[14:16], 32) // bpp + pix := dib[40:] + // 视觉顶部 = DIB 最后一行:红、绿 + off := (h - 1) * w * 4 + binary.LittleEndian.PutUint32(pix[off:off+4], 0x00FF0000) // BGRA → 红 + binary.LittleEndian.PutUint32(pix[off+4:off+8], 0x0000FF00) // 绿 + // 视觉底部 = DIB 首行:蓝、白 + binary.LittleEndian.PutUint32(pix[0:4], 0x000000FF) + binary.LittleEndian.PutUint32(pix[4:8], 0x00FFFFFF) + return dib +} + +// capsPDU 构造 CB_CLIP_CAPS(声明长格式名),handler 需先收到它才会按长名解析 +func capsPDU() []byte { + body := make([]byte, 0, 20) + body = binary.LittleEndian.AppendUint16(body, 1) // cCapSets + body = binary.LittleEndian.AppendUint16(body, 0) // pad + body = binary.LittleEndian.AppendUint16(body, cliprdr.CB_CAPSTYPE_GENERAL) + body = binary.LittleEndian.AppendUint16(body, cliprdr.CB_CAPSTYPE_GENERAL_LEN) + body = binary.LittleEndian.AppendUint32(body, cliprdr.CB_CAPS_VERSION_2) + body = binary.LittleEndian.AppendUint32(body, cliprdr.CB_USE_LONG_FORMAT_NAMES) + p := pduHeader(cliprdr.CB_CLIP_CAPS, 0, len(body)) + return append(p, body...) +} + +func TestDIBToPNG(t *testing.T) { + pngBytes, err := cliprdr.DIBToPNG(fakeDIB()) + if err != nil { + t.Fatalf("DIBToPNG: %v", err) + } + img, err := png.Decode(bytes.NewReader(pngBytes)) + if err != nil { + t.Fatalf("decoded PNG invalid: %v", err) + } + if img.Bounds().Dx() != 2 || img.Bounds().Dy() != 2 { + t.Fatalf("unexpected size %v", img.Bounds()) + } + r, g, b, _ := img.At(0, 0).RGBA() // RGBA() 分量为 16 位 + if r != 0xFFFF || g != 0 || b != 0 { + t.Fatalf("pixel(0,0) want red, got r=%x g=%x b=%x", r, g, b) + } +} + +func TestFormatDataResponseDIBConvertedToPNG(t *testing.T) { + h := cliprdr.NewHandler(nil, nil) + sender := &captureSender{} + h.Sender(sender) + + var gotPNG []byte + h.SetImageCallbacks(func(p []byte) { gotPNG = p }, nil) + + // 先声明能力(长格式名),再让服务端提供仅含 CF_DIB 的 FormatList + h.Process(capsPDU()) + h.Process(mkFormatListLong([]cliprdr.CliprdrFormat{{FormatId: cliprdr.CF_DIB}})) + h.Process(responsePDU(cliprdr.CF_DIB, fakeDIB())) + + if len(gotPNG) == 0 { + t.Fatal("expected PNG bytes via onRemoteClipboardImage") + } + if _, err := png.Decode(bytes.NewReader(gotPNG)); err != nil { + t.Fatalf("converted PNG invalid: %v", err) + } +} + +func TestFormatListPNGTriggersRequest(t *testing.T) { + h := cliprdr.NewHandler(nil, nil) + sender := &captureSender{} + h.Sender(sender) + + h.Process(capsPDU()) + h.Process(mkFormatListLong([]cliprdr.CliprdrFormat{ + {FormatId: 0xC3F0, FormatName: "PNG"}, + {FormatId: cliprdr.CF_UNICODETEXT, FormatName: ""}, + })) + + requests := sender.findByType(0x0004) // CB_FORMAT_DATA_REQUEST + for _, p := range requests { + id := binary.LittleEndian.Uint32(p[8:12]) + if id == 0xC3F0 { // 应按服务端给出的 PNG 注册 ID 发起请求 + return + } + } + t.Fatal("expected PNG format data request for server-offered PNG format") +} + +func TestFormatDataRequestPNGServedFromProvider(t *testing.T) { + h := cliprdr.NewHandler(nil, nil) + sender := &captureSender{} + h.Sender(sender) + h.SetImageCallbacks(nil, func() []byte { return []byte("FAKEPNG") }) + + h.Process(requestPDU(cliprdr.CF_PNG)) + + responses := sender.findByType(0x0005) // CB_FORMAT_DATA_RESPONSE + for _, p := range responses { + if binary.LittleEndian.Uint16(p[2:]) == 1 && bytes.Equal(p[8:15], []byte("FAKEPNG")) { + return + } + } + t.Fatalf("expected PNG data served from provider, got %d responses", len(responses)) +} + +// --- helpers: 构造标准 CLIPRDR PDU --- + +func pduHeader(msgType, msgFlags uint16, bodyLen int) []byte { + h := make([]byte, 8) + binary.LittleEndian.PutUint16(h[0:], msgType) + binary.LittleEndian.PutUint16(h[2:], msgFlags) + binary.LittleEndian.PutUint32(h[4:], uint32(bodyLen)) + return h +} + +func requestPDU(formatId uint32) []byte { + p := pduHeader(0x0004, 0, 4) + return binary.LittleEndian.AppendUint32(p, formatId) +} + +func responsePDU(formatId uint32, data []byte) []byte { + _ = formatId // 响应 PDU 本身不含格式 ID,格式由我方请求记录 + p := pduHeader(0x0005, 1, len(data)) + return append(p, data...) +} + +func mkFormatListLong(formats []cliprdr.CliprdrFormat) []byte { + var body bytes.Buffer + for _, f := range formats { + binary.Write(&body, binary.LittleEndian, f.FormatId) + body.Write(encodeUTF16(f.FormatName)) + body.Write([]byte{0, 0}) + } + p := pduHeader(0x0002, 0, body.Len()) + return append(p, body.Bytes()...) +} + +func encodeUTF16(s string) []byte { + u := []uint16{} + for _, r := range s { + u = append(u, uint16(r)) + } + out := make([]byte, len(u)*2) + for i, v := range u { + binary.LittleEndian.PutUint16(out[i*2:], v) + } + return out +} diff --git a/plugin/cliprdr/cliprdr_test.go b/plugin/cliprdr/cliprdr_test.go new file mode 100644 index 0000000..267358f --- /dev/null +++ b/plugin/cliprdr/cliprdr_test.go @@ -0,0 +1,44 @@ +// cliprdr_test.go +package cliprdr_test + +import ( + "fmt" + "testing" + + "git.zeroonesoft.cn/golib/rdplib/plugin/cliprdr" +) + +func TestClipboardText(t *testing.T) { + // Test text clipboard functionality + testText := "Hello, Clipboard!" + cliprdr.SetClipboardText(testText) + + result := cliprdr.GetClipboardText() + if result != testText { + t.Errorf("Expected %s, got %s", testText, result) + } + fmt.Printf("Text clipboard test passed: %s\n", result) +} + +func TestFormatList(t *testing.T) { + // Test format list (should contain only text format) + formats := cliprdr.GetFormatList() + if len(formats) == 0 { + t.Error("Format list should not be empty") + } + fmt.Printf("Formats: %v\n", formats) +} + +func TestOpenClipboard(t *testing.T) { + // Test OpenClipboard (no-op in generic mode) + ok := cliprdr.OpenClipboard(0) + if !ok { + t.Error("OpenClipboard should return true in generic mode") + } + ok = cliprdr.CloseClipboard() + if !ok { + t.Error("CloseClipboard should return true in generic mode") + } + fmt.Println("OpenClipboard test passed") +} + diff --git a/plugin/cliprdr/cliprdr_types.go b/plugin/cliprdr/cliprdr_types.go new file mode 100644 index 0000000..64ac90c --- /dev/null +++ b/plugin/cliprdr/cliprdr_types.go @@ -0,0 +1,327 @@ +package cliprdr + +import ( + "bytes" + + "github.com/lunixbochs/struc" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/plugin" +) + +/** + * Initialization Sequence\n + * Client Server\n + * | |\n + * |<----------------------Server Clipboard Capabilities PDU-----------------|\n + * |<-----------------------------Monitor Ready PDU--------------------------|\n + * |-----------------------Client Clipboard Capabilities PDU---------------->|\n + * |---------------------------Temporary Directory PDU---------------------->|\n + * |-------------------------------Format List PDU-------------------------->|\n + * |<--------------------------Format List Response PDU----------------------|\n + * + */ + +/** + * Data Transfer Sequences\n + * Shared Local\n + * Clipboard Owner Clipboard Owner\n + * | |\n + * |-------------------------------------------------------------------------|\n _ + * |-------------------------------Format List PDU-------------------------->|\n | + * |<--------------------------Format List Response PDU----------------------|\n _| Copy + * Sequence + * |<---------------------Lock Clipboard Data PDU (Optional)-----------------|\n + * |-------------------------------------------------------------------------|\n + * |-------------------------------------------------------------------------|\n _ + * |<--------------------------Format Data Request PDU-----------------------|\n | Paste + * Sequence Palette, + * |---------------------------Format Data Response PDU--------------------->|\n _| Metafile, + * File List Data + * |-------------------------------------------------------------------------|\n + * |-------------------------------------------------------------------------|\n _ + * |<------------------------Format Contents Request PDU---------------------|\n | Paste + * Sequence + * |-------------------------Format Contents Response PDU------------------->|\n _| File + * Stream Data + * |<---------------------Lock Clipboard Data PDU (Optional)-----------------|\n + * |-------------------------------------------------------------------------|\n + * + */ + +const ( + ChannelName = plugin.CLIPRDR_SVC_CHANNEL_NAME + ChannelOption = plugin.CHANNEL_OPTION_INITIALIZED | plugin.CHANNEL_OPTION_ENCRYPT_RDP | + plugin.CHANNEL_OPTION_COMPRESS_RDP | plugin.CHANNEL_OPTION_SHOW_PROTOCOL +) + +type MsgType uint16 + +const ( + CB_MONITOR_READY = 0x0001 + CB_FORMAT_LIST = 0x0002 + CB_FORMAT_LIST_RESPONSE = 0x0003 + CB_FORMAT_DATA_REQUEST = 0x0004 + CB_FORMAT_DATA_RESPONSE = 0x0005 + CB_TEMP_DIRECTORY = 0x0006 + CB_CLIP_CAPS = 0x0007 + CB_FILECONTENTS_REQUEST = 0x0008 + CB_FILECONTENTS_RESPONSE = 0x0009 + CB_LOCK_CLIPDATA = 0x000A + CB_UNLOCK_CLIPDATA = 0x000B +) + +type MsgFlags uint16 + +const ( + CB_RESPONSE_OK = 0x0001 + CB_RESPONSE_FAIL = 0x0002 + CB_ASCII_NAMES = 0x0004 +) + +type DwFlags uint32 + +const ( + FILECONTENTS_SIZE = 0x00000001 + FILECONTENTS_RANGE = 0x00000002 +) + +type CliprdrPDUHeader struct { + MsgType uint16 `struc:"little"` + MsgFlags uint16 `struc:"little"` + DataLen uint32 `struc:"little"` +} + +func NewCliprdrPDUHeader(mType, flags uint16, ln uint32) *CliprdrPDUHeader { + return &CliprdrPDUHeader{ + MsgType: mType, + MsgFlags: flags, + DataLen: ln, + } +} +func (h *CliprdrPDUHeader) serialize() []byte { + b := &bytes.Buffer{} + core.WriteUInt16LE(h.MsgType, b) + core.WriteUInt16LE(h.MsgFlags, b) + core.WriteUInt32LE(h.DataLen, b) + return b.Bytes() +} + +type CliprdrGeneralCapabilitySet struct { + CapabilitySetType uint16 `struc:"little"` + CapabilitySetLength uint16 `struc:"little"` + Version uint32 `struc:"little"` + GeneralFlags uint32 `struc:"little"` +} + +const ( + CB_CAPSTYPE_GENERAL = 0x0001 +) + +type CliprdrCapabilitySets struct { + CapabilitySetType uint16 `struc:"little"` + LengthCapability uint16 `struc:"little"` + Version uint32 `struc:"little"` + GeneralFlags uint32 `struc:"little"` +} +type CliprdrCapabilitiesPDU struct { + CCapabilitiesSets uint16 `struc:"little,sizeof=CapabilitySets"` + Pad1 uint16 `struc:"little"` + CapabilitySets []CliprdrGeneralCapabilitySet `struc:"little"` +} + +type CliprdrMonitorReady struct { +} + +type GeneralFlags uint32 + +const ( + CB_USE_LONG_FORMAT_NAMES = 0x00000002 + CB_STREAM_FILECLIP_ENABLED = 0x00000004 + CB_FILECLIP_NO_FILE_PATHS = 0x00000008 + CB_CAN_LOCK_CLIPDATA = 0x00000010 + CB_HUGE_FILE_SUPPORT_ENABLED = 0x00000020 +) + +const ( + CB_CAPS_VERSION_1 = 0x00000001 + CB_CAPS_VERSION_2 = 0x00000002 +) +const ( + CB_CAPSTYPE_GENERAL_LEN = 12 +) + +const ( + FD_CLSID = 0x00000001 + FD_SIZEPOINT = 0x00000002 + FD_ATTRIBUTES = 0x00000004 + FD_CREATETIME = 0x00000008 + FD_ACCESSTIME = 0x00000010 + FD_WRITESTIME = 0x00000020 + FD_FILESIZE = 0x00000040 + FD_PROGRESSUI = 0x00004000 + FD_LINKUI = 0x00008000 +) + +const ( + FILE_ATTRIBUTE_DIRECTORY = 0x00000010 + FILE_ATTRIBUTE_ARCHIVE = 0x00000020 +) + +// fileDescriptorZero is a reusable zero block for FileDescriptor padding writes. +var fileDescriptorZero [512]byte + +type FileGroupDescriptor struct { + CItems uint32 `struc:"little"` + Fgd []FileDescriptor `struc:"sizefrom=CItems"` +} +type FileDescriptor struct { + Flags uint32 `struc:"little"` + Clsid [16]byte `struc:"little"` + Sizel [8]byte `struc:"little"` + Pointl [8]byte `struc:"little"` + FileAttributes uint32 `struc:"little"` + CreationTime [8]byte `struc:"little"` + LastAccessTime [8]byte `struc:"little"` + LastWriteTime []byte `struc:"[8]byte"` //8 + FileSizeHigh uint32 `struc:"little"` + FileSizeLow uint32 `struc:"little"` + FileName []byte `struc:"[512]byte"` +} + +func (f *FileGroupDescriptor) Unpack(b []byte) error { + r := bytes.NewReader(b) + return struc.Unpack(r, f) +} + +func (f *FileDescriptor) serialize() []byte { + b := &bytes.Buffer{} + core.WriteUInt32LE(f.Flags, b) + b.Write(fileDescriptorZero[:32]) + core.WriteUInt32LE(f.FileAttributes, b) + b.Write(fileDescriptorZero[:16]) + core.WriteBytes(f.LastWriteTime[:], b) + core.WriteUInt32LE(f.FileSizeHigh, b) + core.WriteUInt32LE(f.FileSizeLow, b) + b.Write(f.FileName) + if pad := 512 - len(f.FileName); pad > 0 { + b.Write(fileDescriptorZero[:pad]) + } + return b.Bytes() +} + +func (f *FileDescriptor) isDir() bool { + if f.Flags&FD_ATTRIBUTES != 0 { + return f.FileAttributes&FILE_ATTRIBUTE_DIRECTORY != 0 + } + return false +} + +func (f *FileDescriptor) hasFileSize() bool { + return f.Flags&FD_FILESIZE != 0 +} + +// temp dir +type CliprdrTempDirectory struct { + SzTempDir []byte `struc:"[260]byte"` +} + +// format list +type CliprdrFormat struct { + FormatId uint32 + FormatName string +} +type CliprdrFormatList struct { + NumFormats uint32 + Formats []CliprdrFormat +} +type ClipboardFormats uint16 + +const ( + CB_FORMAT_HTML = 0xD010 + CB_FORMAT_PNG = 0xD011 + CB_FORMAT_JPEG = 0xD012 + CB_FORMAT_GIF = 0xD013 + CB_FORMAT_TEXTURILIST = 0xD014 + CB_FORMAT_GNOMECOPIEDFILES = 0xD015 + CB_FORMAT_MATECOPIEDFILES = 0xD016 +) + +// Standard clipboard format IDs +const ( + CF_TEXT = 1 + CF_DIB = 8 + CF_UNICODETEXT = 13 + CF_DIBV5 = 17 + // CF_PNG 是客户端自定义的注册格式 ID。注册格式在线上按"名称"匹配, + // Windows 侧对应 RegisterClipboardFormat("PNG"),ID 仅需 >0xC000 且客户端自定。 + CF_PNG = 0xC0BC + // CF_HTML_FORMAT_ID 是客户端自定的 "HTML Format" 注册格式 ID(同上)。 + CF_HTML_FORMAT_ID = 0xC0BE +) + +// FormatNamePNG 是 Windows 剪贴板的 PNG 注册格式名 +const FormatNamePNG = "PNG" + +// FormatNameHTML 是 Windows 剪贴板的 HTML 注册格式名(MS-Doc: HTML Clipboard Format) +const FormatNameHTML = "HTML Format" + +// lock or unlock +type CliprdrCtrlClipboardData struct { + ClipDataId uint32 +} + +// format data +type CliprdrFormatDataRequest struct { + RequestedFormatId uint32 +} +type CliprdrFormatDataResponse struct { + RequestedFormatData []byte +} + +// file contents +type CliprdrFileContentsRequest struct { + StreamId uint32 `struc:"little"` + Lindex uint32 `struc:"little"` + DwFlags uint32 `struc:"little"` + NPositionLow uint32 `struc:"little"` + NPositionHigh uint32 `struc:"little"` + CbRequested uint32 `struc:"little"` + ClipDataId uint32 `struc:"little"` +} + +func FileContentsSizeRequest(i uint32) *CliprdrFileContentsRequest { + return &CliprdrFileContentsRequest{ + StreamId: 1, + Lindex: i, + DwFlags: FILECONTENTS_SIZE, + NPositionLow: 0, + NPositionHigh: 0, + CbRequested: 65535, + ClipDataId: 0, + } +} + +type CliprdrFileContentsResponse struct { + StreamId uint32 + CbRequested uint32 + RequestedData []byte +} + +func (resp *CliprdrFileContentsResponse) Unpack(b []byte) { + r := bytes.NewReader(b) + resp.StreamId, _ = core.ReadUInt32LE(r) + resp.CbRequested = uint32(r.Len()) + resp.RequestedData, _ = core.ReadBytes(int(resp.CbRequested), r) +} + +// sendClipPDU sends a CLIPRDR PDU with the standard 8-byte header +// (msgType + msgFlags + dataLen) prepended to body. +func sendClipPDU(sender core.ChannelSender, msgType, msgFlags uint16, body []byte) { + b := &bytes.Buffer{} + core.WriteUInt16LE(msgType, b) + core.WriteUInt16LE(msgFlags, b) + core.WriteUInt32LE(uint32(len(body)), b) + b.Write(body) + sender.SendToChannel(ChannelName, b.Bytes()) +} diff --git a/plugin/cliprdr/file_clip.go b/plugin/cliprdr/file_clip.go new file mode 100644 index 0000000..9d629a8 --- /dev/null +++ b/plugin/cliprdr/file_clip.go @@ -0,0 +1,479 @@ +package cliprdr + +// file_clip.go implements the MS-RDPECLIP file transfer sequences +// (§3.1.5.4.4 Copy File Sequence / §3.1.5.4.5 File Transfer Data Sequence): +// +// - Server → client (files copied in the remote session): the server's +// Format List contains CF_HDROP; we fetch the path list via a Format +// Data Request, then each file's bytes via FileContentsRequest +// (FILECONTENTS_SIZE, then FILECONTENTS_RANGE chunks). +// - Client → server (files staged by the local UI): staged files are +// advertised as CF_HDROP in our Format List; the server fetches the +// DROPFILES path list via Format Data Request, then file bytes via +// FileContentsRequest, which we answer from the staged buffers. + +import ( + "bytes" + "encoding/binary" + "errors" + "log/slog" + "unicode/utf16" +) + +// CF_HDROP 是标准剪贴板格式 15(文件列表)。Windows 在 short/long 两种 +// 格式名模式下都按 ID 识别标准格式,名称可留空。 +const CF_HDROP = 15 + +// FormatNameHDrop 是部分实现(FreeRDP 等)在 short name 里使用的名称。 +const FormatNameHDrop = "Hdrop" + +// CF_DROP_EFFECT 是客户端自定的 "Preferred DropEffect" 注册格式 ID +// (>0xC000 即可,线上按名称匹配)。Windows 资源管理器判定文件粘贴是否 +// 可用时查询该格式;缺失或 FAIL 会导致右键菜单"粘贴"置灰。 +const CF_DROP_EFFECT = 0xC0C0 + +// FormatNameDropEffect 是 Windows 的 DropEffect 注册格式名。 +const FormatNameDropEffect = "Preferred DropEffect" + +// DROPEFFECT_COPY 是 DropEffect 值:我们的文件提供方式是复制。 +const DROPEFFECT_COPY = 1 + +// MS-RDPECLIP 文件传输采用 FileGroupDescriptorW + FileContents 注册格式对 +// (配合 CB_FILECLIP_NO_FILE_PATHS 能力位,mstsc 同款)。ID 客户端自定, +// 服务器按我们的 Format List 中通告的 ID 回查。 +const ( + CF_FILE_GROUP_DESCRIPTORW = 0xC0C6 + CF_FILE_CONTENTS = 0xC0C7 +) + +// FormatNameFileGroupDescriptorW / FormatNameFileContents 是线上的注册格式名。 +const ( + FormatNameFileGroupDescriptorW = "FileGroupDescriptorW" + FormatNameFileContents = "FileContents" +) + +// fileDescriptorSize 是 FILE_DESCRIPTOR 的线上字节数 +// (flags4 + clsid16 + sizl8 + pointl8 + attr4 + 3×time8 + sizeHi4 + sizeLo4 + name520)。 +const fileDescriptorSize = 592 + +// fileDescriptorNameBytes 是 cFileName 字段长度(260 WCHAR = 520 字节)。 +const fileDescriptorNameBytes = 520 + +// dropfilesHeaderLen 是 DROPFILES 结构长度(pFiles 4 + pt 8 + fNC 4 + fWide 4)。 +const dropfilesHeaderLen = 20 + +// fileRangeChunk 是我方主动拉取远端文件时的每段 RANGE 大小。服务器可能 +// 裁剪返回长度,接收端按"追加到累计长度"推进,不假设服务器返回整段。 +const fileRangeChunk = 256 * 1024 + +// maxDescriptorItems 防御性上限:一份描述符列表最多 65536 项(约 38MB +// 载荷),超出视为协议错位。远端复制上万文件的目录仍被接受。 +const maxDescriptorItems = 65536 + +// maxFileTotalBytes 拒绝异常大的 SIZE 声明(>4GB 意味着协议解析已错位)。 +const maxFileTotalBytes = int64(4) << 30 + +// LocalFile is a file staged by the local UI for client → server transfer. +type LocalFile struct { + Name string + Data []byte +} + +// fileStream tracks one in-flight remote → local transfer (our StreamId). +type fileStream struct { + index int + name string + total int64 + got int64 + buf []byte + sizing bool // true until the FILECONTENTS_SIZE response arrives +} + +// --- Public API ------------------------------------------------------------- + +// SetFileCallbacks attaches file clipboard callbacks (all optional): +// +// - onRemoteFiles is called with the file names when the server's +// clipboard holds files (CF_HDROP path list received). +// - onFileData is called with the assembled bytes of one file after +// RequestRemoteFile finishes; data is nil when the transfer failed. +// - onFileProgress is called after every received chunk. +func (h *CliprdrHandler) SetFileCallbacks( + onRemoteFiles func(names []string), + onFileData func(index int, name string, data []byte), + onFileProgress func(index int, received, total int64), +) { + h.onRemoteFiles = onRemoteFiles + h.onFileData = onFileData + h.onFileProgress = onFileProgress +} + +// SetLocalFiles stages files as the local (client) clipboard file content and +// advertises them to the server as CF_HDROP. +func (h *CliprdrHandler) SetLocalFiles(files []LocalFile) { + h.mu.Lock() + h.localFiles = files + h.mu.Unlock() + // 暂存文件是真实的本地剪贴板变更(UI 主动动作),直接重发 Format List, + // 不走 suppressNextLocalChange(它属于远端回显抑制语义)。 + if h.channelSender != nil { + h.sendFormatList() + slog.Debug("cliprdr: local files staged, sent Format List", "files", len(files)) + } +} + +// buildDropEffectReply 返回 "Preferred DropEffect" 的应答载荷 +// (4 字节 DWORD DROPEFFECT_COPY)。 +func buildDropEffectReply() []byte { + b := make([]byte, 4) + binary.LittleEndian.PutUint32(b, DROPEFFECT_COPY) + return b +} + +// BuildFileGroupDescriptorW 将暂存文件编码为 FILE_GROUP_DESCRIPTORW 载荷 +// (MS-RDPECLIP 2.2.5.2.3.1):cItems + FILE_DESCRIPTOR[cItems]。 +func BuildFileGroupDescriptorW(files []LocalFile) []byte { + b := &bytes.Buffer{} + u32 := func(v uint32) { binary.Write(b, binary.LittleEndian, v) } + u32(uint32(len(files))) + for _, f := range files { + u32(FD_ATTRIBUTES | FD_FILESIZE | FD_PROGRESSUI) + b.Write(make([]byte, 16)) // clsid + b.Write(make([]byte, 8)) // sizl + b.Write(make([]byte, 8)) // pointl + u32(FILE_ATTRIBUTE_ARCHIVE) + b.Write(make([]byte, 24)) // creation/access/write time(未设时间标志,值无效) + u32(uint32(uint64(len(f.Data)) >> 32)) + u32(uint32(len(f.Data))) + name := make([]byte, fileDescriptorNameBytes) + u16 := utf16.Encode([]rune(f.Name)) + for i := 0; i < len(u16) && 2*i+1 < fileDescriptorNameBytes-1; i++ { + binary.LittleEndian.PutUint16(name[2*i:], u16[i]) + } + b.Write(name) + } + return b.Bytes() +} + +// ParseFileGroupDescriptorW 解码服务器发来的 FILE_GROUP_DESCRIPTORW, +// 返回文件名与文件大小。 +func ParseFileGroupDescriptorW(body []byte) ([]string, []int64, error) { + if len(body) < 4 { + return nil, nil, errors.New("cliprdr: FileGroupDescriptorW too short") + } + cItems := binary.LittleEndian.Uint32(body[0:]) + // 上限只做溢出/内存炸弹防护(uint32 项数 × 592B 的乘积合法性), + // 实际约束是载荷长度必须恰好容纳 cItems 个描述符——远端复制含 + // 数千文件的目录是正常操作,不能按固定小数目拒绝。 + if cItems > maxDescriptorItems || len(body) < 4+int(cItems)*fileDescriptorSize { + return nil, nil, errors.New("cliprdr: bad FileGroupDescriptorW cItems/length") + } + names := make([]string, 0, cItems) + sizes := make([]int64, 0, cItems) + for i := 0; i < int(cItems); i++ { + d := body[4+i*fileDescriptorSize : (i+1)*fileDescriptorSize] + flags := binary.LittleEndian.Uint32(d[0:]) + size := uint64(binary.LittleEndian.Uint32(d[68:])) | + uint64(binary.LittleEndian.Uint32(d[64:]))<<32 + rawName := d[72:] + // 260 WCHAR UTF-16LE,NUL 结尾 + end := 0 + for end+1 < len(rawName) { + if rawName[end] == 0 && rawName[end+1] == 0 { + break + } + end += 2 + } + name := decodeUTF16LE(rawName[:end]) + if flags&FD_FILESIZE == 0 { + size = ^uint64(0) // 大小未知,由 FileContentsRequest(SIZE) 探测 + } + names = append(names, name) + sizes = append(sizes, int64(size)) + } + return names, sizes, nil +} + +// ClearLocalFiles removes staged files and re-advertises the Format List. +func (h *CliprdrHandler) ClearLocalFiles() { + h.mu.Lock() + had := len(h.localFiles) > 0 + h.localFiles = nil + h.mu.Unlock() + if had && h.channelSender != nil { + h.sendFormatList() + } +} + +// RemoteFileNames returns the paths from the server's last CF_HDROP payload. +func (h *CliprdrHandler) RemoteFileNames() []string { + h.mu.Lock() + defer h.mu.Unlock() + out := make([]string, len(h.remoteFiles)) + copy(out, h.remoteFiles) + return out +} + +// RequestRemoteFile starts downloading file `index` (order within the last +// file list) from the server. Progress goes to onFileProgress; the +// assembled bytes go to onFileData(index, name, data) — data is nil on failure. +func (h *CliprdrHandler) RequestRemoteFile(index int) error { + h.mu.Lock() + if index < 0 || index >= len(h.remoteFiles) { + h.mu.Unlock() + return errors.New("cliprdr: file index out of range") + } + name := h.remoteFiles[index] + size := int64(-1) + if index < len(h.remoteFileSizes) { + size = h.remoteFileSizes[index] + } + h.nextStreamID++ + sid := h.nextStreamID + h.mu.Unlock() + + if size >= 0 { + // FileGroupDescriptorW 已声明大小:直接从 0 开始拉 RANGE + h.mu.Lock() + h.streams[sid] = &fileStream{index: index, name: name, total: size} + h.mu.Unlock() + if size == 0 { // 零字节文件:无 RANGE 可发,直接完成 + h.processFileContentsResponse(u32le(sid), CB_RESPONSE_OK) + return nil + } + h.sendFileContentsRequest(sid, uint32(index), FILECONTENTS_RANGE, 0, fileRangeChunk) + slog.Debug("cliprdr: requesting remote file", "index", index, "name", name, + "streamId", sid, "size", size) + return nil + } + + h.mu.Lock() + h.streams[sid] = &fileStream{index: index, name: name, sizing: true} + h.mu.Unlock() + h.sendFileContentsRequest(sid, uint32(index), FILECONTENTS_SIZE, 0, 8) + slog.Debug("cliprdr: requesting remote file", "index", index, "name", name, "streamId", sid) + return nil +} + +// --- Wire helpers ----------------------------------------------------------- + +// sendFileContentsRequest sends CB_FILECONTENTS_REQUEST (MS-RDPECLIP 2.2.5.2.3). +func (h *CliprdrHandler) sendFileContentsRequest(streamId, lindex, dwFlags, posLow, cbRequested uint32) { + b := make([]byte, 24) + binary.LittleEndian.PutUint32(b[0:], streamId) + binary.LittleEndian.PutUint32(b[4:], lindex) + binary.LittleEndian.PutUint32(b[8:], dwFlags) + binary.LittleEndian.PutUint32(b[12:], posLow) + binary.LittleEndian.PutUint32(b[16:], 0) // NPositionHigh:4GB 内恒 0 + binary.LittleEndian.PutUint32(b[20:], cbRequested) + h.sendPDU(CB_FILECONTENTS_REQUEST, 0, b) +} + +// processFileContentsRequest answers the server's request for staged file +// bytes (client → server direction). Failure replies carry only StreamId. +func (h *CliprdrHandler) processFileContentsRequest(body []byte) { + if len(body) < 24 { + slog.Warn("cliprdr: short FileContentsRequest", "len", len(body)) + return + } + streamId := binary.LittleEndian.Uint32(body[0:]) + lindex := binary.LittleEndian.Uint32(body[4:]) + dwFlags := binary.LittleEndian.Uint32(body[8:]) + pos := uint64(binary.LittleEndian.Uint32(body[12:])) | uint64(binary.LittleEndian.Uint32(body[16:]))<<32 + cbRequested := binary.LittleEndian.Uint32(body[20:]) + // body[24:28] 是 ClipDataId,仅在 CB_CAN_LOCK_CLIPDATA 协商后出现,此处忽略。 + + h.mu.Lock() + files := h.localFiles + h.mu.Unlock() + + var data []byte + if lindex < uint32(len(files)) { + data = files[lindex].Data + } + if data == nil { + slog.Warn("cliprdr: FileContentsRequest for unknown file", "lindex", lindex) + h.sendPDU(CB_FILECONTENTS_RESPONSE, CB_RESPONSE_FAIL, u32le(streamId)) + return + } + + var payload []byte + switch { + case dwFlags&FILECONTENTS_SIZE != 0: + payload = make([]byte, 8) + binary.LittleEndian.PutUint64(payload, uint64(len(data))) + case dwFlags&FILECONTENTS_RANGE != 0: + start := pos + end := pos + uint64(cbRequested) + if end > uint64(len(data)) || end < start { + end = uint64(len(data)) + } + if start < uint64(len(data)) { + payload = data[start:end] + } + default: + slog.Warn("cliprdr: FileContentsRequest without SIZE/RANGE", "dwFlags", dwFlags) + h.sendPDU(CB_FILECONTENTS_RESPONSE, CB_RESPONSE_FAIL, u32le(streamId)) + return + } + + out := make([]byte, 4+len(payload)) + binary.LittleEndian.PutUint32(out, streamId) + copy(out[4:], payload) + h.sendPDU(CB_FILECONTENTS_RESPONSE, CB_RESPONSE_OK, out) + slog.Debug("cliprdr: served file contents", "lindex", lindex, "flags", dwFlags, + "pos", pos, "len", len(payload)) +} + +// processFileContentsResponse consumes the server's reply for one of our +// outstanding RequestRemoteFile transfers, chaining RANGE requests until the +// declared size has been received. +func (h *CliprdrHandler) processFileContentsResponse(body []byte, msgFlags uint16) { + if len(body) < 4 { + return + } + streamId := binary.LittleEndian.Uint32(body[0:]) + data := body[4:] + + h.mu.Lock() + st, ok := h.streams[streamId] + if !ok { + h.mu.Unlock() + slog.Debug("cliprdr: FileContentsResponse for unknown stream", "streamId", streamId) + return + } + fail := msgFlags&CB_RESPONSE_OK == 0 + if fail { + delete(h.streams, streamId) + } else if st.sizing { + // SIZE 响应固定 8 字节 uint64(部分实现只回 4 字节,向下兼容) + if len(data) < 4 { + delete(h.streams, streamId) + fail = true + } else { + var size uint64 + if len(data) >= 8 { + size = binary.LittleEndian.Uint64(data) + } else { + size = uint64(binary.LittleEndian.Uint32(data)) + } + if size > uint64(maxFileTotalBytes) { + slog.Warn("cliprdr: remote file size absurd", "size", size) + delete(h.streams, streamId) + fail = true + } else { + st.sizing = false + st.total = int64(size) + st.buf = make([]byte, 0, size) + } + } + } else { + // RANGE 响应:按实际到达长度追加(服务器允许裁剪返回段) + st.buf = append(st.buf, data...) + st.got = int64(len(st.buf)) + } + index, name, got, total := st.index, st.name, st.got, st.total + requestNext := false + var nextPos uint32 + var finished []byte + if !fail && !st.sizing { + if st.got < st.total { + requestNext = true + nextPos = uint32(st.got) // 4GB 内;超出已由 maxFileTotalBytes 拒绝 + } else { + // 传输完成(含零字节文件:SIZE 后即完成)——锁内取走缓冲防竞态 + finished = make([]byte, total) + copy(finished, st.buf) + delete(h.streams, streamId) + } + } + h.mu.Unlock() + + if fail { + slog.Warn("cliprdr: file transfer failed", "name", name, "streamId", streamId) + if h.onFileData != nil { + h.onFileData(index, name, nil) + } + return + } + + if h.onFileProgress != nil { + h.onFileProgress(index, got, total) + } + + if requestNext { + h.sendFileContentsRequest(streamId, uint32(index), FILECONTENTS_RANGE, nextPos, fileRangeChunk) + return + } + + if finished != nil && h.onFileData != nil { + slog.Debug("cliprdr: file transfer complete", "name", name, "bytes", total) + h.onFileData(index, name, finished) + } +} + +// --- DROPFILES (CF_HDROP format data, MS-RDPECLIP 2.2.5.2.4) ---------------- + +// ParseDropfiles decodes a CF_HDROP payload (DROPFILES header + null- +// terminated path list, UTF-16LE when fWide==1, ASCII otherwise) into paths. +func ParseDropfiles(body []byte) ([]string, error) { + if len(body) < dropfilesHeaderLen { + return nil, errors.New("cliprdr: DROPFILES body too short") + } + pFiles := binary.LittleEndian.Uint32(body[0:]) + fWide := binary.LittleEndian.Uint32(body[16:]) + if pFiles < dropfilesHeaderLen || int(pFiles) > len(body) { + return nil, errors.New("cliprdr: bad DROPFILES pFiles offset") + } + list := body[pFiles:] + var names []string + if fWide == 1 { + // 双 NUL 结尾的 UTF-16LE 路径列表 + start := 0 + for i := 0; i+1 < len(list); i += 2 { + if list[i] == 0 && list[i+1] == 0 { + if i == start { // 连续两个 NUL = 列表结束 + break + } + names = append(names, decodeUTF16LE(list[start:i])) + start = i + 2 + } + } + } else { + start := 0 + for i := 0; i < len(list); i++ { + if list[i] == 0 { + if i == start { + break + } + names = append(names, string(list[start:i])) + start = i + 1 + } + } + } + return names, nil +} + +// BuildDropfiles encodes names into a CF_HDROP payload (UTF-16LE, fWide=1). +func BuildDropfiles(names []string) []byte { + b := &bytes.Buffer{} + u32 := func(v uint32) { binary.Write(b, binary.LittleEndian, v) } + u32(dropfilesHeaderLen) // pFiles + u32(0) // pt.x + u32(0) // pt.y + u32(0) // fNC + u32(1) // fWide = Unicode + for _, n := range names { + b.Write(encodeUTF16LE(n)) + b.Write([]byte{0, 0}) // 路径 NUL 结尾 + } + b.Write([]byte{0, 0}) // 列表终止双 NUL + return b.Bytes() +} + +// u32le returns v as a 4-byte little-endian slice. +func u32le(v uint32) []byte { + b := make([]byte, 4) + binary.LittleEndian.PutUint32(b, v) + return b +} diff --git a/plugin/cliprdr/file_clip_test.go b/plugin/cliprdr/file_clip_test.go new file mode 100644 index 0000000..f834c3f --- /dev/null +++ b/plugin/cliprdr/file_clip_test.go @@ -0,0 +1,527 @@ +package cliprdr + +import ( + "bytes" + "encoding/binary" + "fmt" + "testing" +) + +// capturedPDU 是 fakeSender 捕获的一条完整 CLIPRDR PDU。 +type capturedPDU struct { + msgType uint16 + flags uint16 + body []byte +} + +// fakeSender 捕获所有发出通道的数据,供断言。 +type fakeSender struct{ pdus []capturedPDU } + +func (f *fakeSender) SendToChannel(_ string, data []byte) (int, error) { + cp := make([]byte, len(data)) + copy(cp, data) + if len(cp) < 8 { + return len(cp), nil + } + f.pdus = append(f.pdus, capturedPDU{ + msgType: binary.LittleEndian.Uint16(cp[0:]), + flags: binary.LittleEndian.Uint16(cp[2:]), + body: cp[8:], + }) + return len(cp), nil +} + +func (f *fakeSender) last() capturedPDU { return f.pdus[len(f.pdus)-1] } + +// newFileTestHandler 构造带 fakeSender 且 serverFileClip=true 的 handler +// (模拟服务器 caps 已声明 CB_STREAM_FILECLIP_ENABLED)。 +func newFileTestHandler(t *testing.T) (*CliprdrHandler, *fakeSender) { + t.Helper() + h := NewHandler(nil, nil) + fs := &fakeSender{} + h.Sender(fs) + caps := &bytes.Buffer{} + binary.Write(caps, binary.LittleEndian, uint16(1)) // cCapSets + binary.Write(caps, binary.LittleEndian, uint16(0)) // pad + binary.Write(caps, binary.LittleEndian, uint16(CB_CAPSTYPE_GENERAL)) + binary.Write(caps, binary.LittleEndian, uint16(12)) + binary.Write(caps, binary.LittleEndian, uint32(CB_CAPS_VERSION_2)) + binary.Write(caps, binary.LittleEndian, uint32(CB_USE_LONG_FORMAT_NAMES|CB_STREAM_FILECLIP_ENABLED)) + h.processClipCaps(caps.Bytes()) + if !h.serverFileClip { + t.Fatal("serverFileClip should be set after caps with CB_STREAM_FILECLIP_ENABLED") + } + return h, fs +} + +func TestBuildParseDropfilesRoundTrip(t *testing.T) { + names := []string{`C:\Users\测试\报告 最终版.docx`, "plain.txt", "数据 - 副本 (2).xlsx"} + b := BuildDropfiles(names) + if binary.LittleEndian.Uint32(b[0:]) != dropfilesHeaderLen { + t.Fatalf("pFiles = %d, want %d", binary.LittleEndian.Uint32(b[0:]), dropfilesHeaderLen) + } + if binary.LittleEndian.Uint32(b[16:]) != 1 { + t.Fatal("fWide should be 1 (Unicode)") + } + got, err := ParseDropfiles(b) + if err != nil { + t.Fatal(err) + } + if len(got) != len(names) { + t.Fatalf("got %d names, want %d: %q", len(got), len(names), got) + } + for i := range names { + if got[i] != names[i] { + t.Fatalf("name[%d] = %q, want %q", i, got[i], names[i]) + } + } +} + +func TestParseDropfilesASCII(t *testing.T) { + // fWide=0 的 ASCII 变体(FreeRDP 老服务器可能使用) + list := []byte("a.txt\x00b.bin\x00\x00") + body := make([]byte, dropfilesHeaderLen+len(list)) + binary.LittleEndian.PutUint32(body[0:], dropfilesHeaderLen) + binary.LittleEndian.PutUint32(body[16:], 0) + copy(body[dropfilesHeaderLen:], list) + got, err := ParseDropfiles(body) + if err != nil { + t.Fatal(err) + } + if len(got) != 2 || got[0] != "a.txt" || got[1] != "b.bin" { + t.Fatalf("got %q", got) + } +} + +func TestParseDropfilesRejectsBad(t *testing.T) { + if _, err := ParseDropfiles(make([]byte, 8)); err == nil { + t.Fatal("short body should fail") + } + bad := make([]byte, dropfilesHeaderLen) + binary.LittleEndian.PutUint32(bad[0:], 1<<20) // pFiles 越界 + if _, err := ParseDropfiles(bad); err == nil { + t.Fatal("out-of-range pFiles should fail") + } +} + +func serverFormatListWithHDrop(t *testing.T, h *CliprdrHandler) { + t.Helper() + // long-name Format List:CF_HDROP(id 15) + CF_UNICODETEXT + b := &bytes.Buffer{} + binary.Write(b, binary.LittleEndian, uint32(CF_HDROP)) + b.Write(encodeUTF16LE("")) // 标准格式空名 + b.Write([]byte{0, 0}) + binary.Write(b, binary.LittleEndian, uint32(CF_UNICODETEXT)) + b.Write(encodeUTF16LE("")) + b.Write([]byte{0, 0}) + h.processFormatList(b.Bytes(), 0) +} + +func TestServerFormatListTriggersHDropRequest(t *testing.T) { + h, fs := newFileTestHandler(t) + serverFormatListWithHDrop(t, h) + // 第一条:Format List Response OK;第二条:Format Data Request(CF_HDROP) + if len(fs.pdus) < 2 { + t.Fatalf("expected FORMAT_LIST_RESPONSE + FORMAT_DATA_REQUEST, got %d", len(fs.pdus)) + } + req := fs.pdus[1] + if req.msgType != CB_FORMAT_DATA_REQUEST { + t.Fatalf("msgType=%#x, want FORMAT_DATA_REQUEST", req.msgType) + } + if id := binary.LittleEndian.Uint32(req.body); id != CF_HDROP { + t.Fatalf("requested format=%d, want %d", id, CF_HDROP) + } + if h.lastRequestedFormat != CF_HDROP { + t.Fatalf("lastRequestedFormat=%d", h.lastRequestedFormat) + } +} + +func TestRemoteFilesReceivedViaFormatDataResponse(t *testing.T) { + h, _ := newFileTestHandler(t) + var gotNames []string + h.SetFileCallbacks(func(names []string) { gotNames = names }, nil, nil) + serverFormatListWithHDrop(t, h) + + payload := BuildDropfiles([]string{`D:\share\a.pdf`, `D:\share\b.pdf`}) + h.processFormatDataResponse(payload, CB_RESPONSE_OK) + if len(gotNames) != 2 || gotNames[0] != `D:\share\a.pdf` { + t.Fatalf("onRemoteFiles got %q", gotNames) + } + if names := h.RemoteFileNames(); len(names) != 2 { + t.Fatalf("RemoteFileNames = %q", names) + } +} + +func TestFormatDataRequestServesStagedFiles(t *testing.T) { + h, fs := newFileTestHandler(t) + h.SetLocalFiles([]LocalFile{{Name: "x.bin", Data: []byte{1, 2, 3}}}) + + req := make([]byte, 4) + binary.LittleEndian.PutUint32(req, CF_FILE_GROUP_DESCRIPTORW) + h.processFormatDataRequest(req) + + resp := fs.last() + if resp.msgType != CB_FORMAT_DATA_RESPONSE || resp.flags != CB_RESPONSE_OK { + t.Fatalf("msgType=%#x flags=%#x", resp.msgType, resp.flags) + } + names, sizes, err := ParseFileGroupDescriptorW(resp.body) + if err != nil || len(names) != 1 || names[0] != "x.bin" || sizes[0] != 3 { + t.Fatalf("FileGroupDescriptorW parse=%q %v err=%v", names, sizes, err) + } +} + +func TestFormatDataRequestNoFilesFails(t *testing.T) { + h, fs := newFileTestHandler(t) + req := make([]byte, 4) + binary.LittleEndian.PutUint32(req, CF_FILE_GROUP_DESCRIPTORW) + h.processFormatDataRequest(req) + if fs.last().flags != CB_RESPONSE_FAIL { + t.Fatal("expected FAIL when no staged files") + } +} + +// TestBuildParseFileGroupDescriptorWRoundTrip:编码→解码保持名称与大小。 +func TestBuildParseFileGroupDescriptorWRoundTrip(t *testing.T) { + files := []LocalFile{ + {Name: "报 告 最终版.docx", Data: make([]byte, 5)}, + {Name: "小.txt", Data: []byte("hi")}, + } + payload := BuildFileGroupDescriptorW(files) + if len(payload) != 4+2*fileDescriptorSize { + t.Fatalf("payload=%d bytes, want %d", len(payload), 4+2*fileDescriptorSize) + } + names, sizes, err := ParseFileGroupDescriptorW(payload) + if err != nil { + t.Fatal(err) + } + if len(names) != 2 { + t.Fatalf("names=%q", names) + } + if names[0] != "报 告 最终版.docx" || names[1] != "小.txt" { + t.Fatalf("names=%q", names) + } + if sizes[0] != 5 || sizes[1] != 2 { + t.Fatalf("sizes=%v", sizes) + } +} + +// TestServerFileDescriptorListFlow:服务器通告 "FileGroupDescriptorW" 时 +// 应以其 ID 请求格式数据,并用描述符填充远端文件列表(含大小)。 +func TestServerFileDescriptorListFlow(t *testing.T) { + h, fs := newFileTestHandler(t) + var gotNames []string + h.SetFileCallbacks(func(names []string) { gotNames = names }, nil, nil) + + // 服务器 Format List:FileGroupDescriptorW(id 0xC181) + FileContents + DropEffect + b := &bytes.Buffer{} + writeNamed := func(id uint32, name string) { + binary.Write(b, binary.LittleEndian, id) + b.Write(encodeUTF16LE(name)) + b.Write([]byte{0, 0}) + } + writeNamed(0xC181, FormatNameFileGroupDescriptorW) + writeNamed(0xC182, FormatNameFileContents) + writeNamed(0xC17E, FormatNameDropEffect) + h.processFormatList(b.Bytes(), 0) + + req := fs.last() + if req.msgType != CB_FORMAT_DATA_REQUEST || + binary.LittleEndian.Uint32(req.body) != 0xC181 { + t.Fatalf("expected FormatDataRequest(0xC181), got %#x id=%d", + req.msgType, binary.LittleEndian.Uint32(req.body)) + } + + payload := BuildFileGroupDescriptorW([]LocalFile{{Name: `C:\doc\a.pdf`, Data: make([]byte, 9)}}) + h.processFormatDataResponse(payload, CB_RESPONSE_OK) + if len(gotNames) != 1 || gotNames[0] != `C:\doc\a.pdf` { + t.Fatalf("onRemoteFiles got %q", gotNames) + } + + // 已知大小:RequestRemoteFile 应直接发 RANGE(0) 而非 SIZE + if err := h.RequestRemoteFile(0); err != nil { + t.Fatal(err) + } + fr := fs.last() + if binary.LittleEndian.Uint32(fr.body[8:]) != FILECONTENTS_RANGE || + binary.LittleEndian.Uint32(fr.body[12:]) != 0 { + t.Fatalf("want direct RANGE(0), flags=%d pos=%d", + binary.LittleEndian.Uint32(fr.body[8:]), binary.LittleEndian.Uint32(fr.body[12:])) + } +} + +// TestDropEffectRequest:资源管理器右键菜单查询 "Preferred DropEffect", +// 有暂存文件时必须回 OK + DROPEFFECT_COPY,否则"粘贴"被置灰。 +func TestDropEffectRequest(t *testing.T) { + h, fs := newFileTestHandler(t) + req := make([]byte, 4) + binary.LittleEndian.PutUint32(req, CF_DROP_EFFECT) + h.processFormatDataRequest(req) + if fs.last().flags != CB_RESPONSE_FAIL { + t.Fatal("no staged files: expected FAIL for DropEffect") + } + + h.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + h.processFormatDataRequest(req) + resp := fs.last() + if resp.msgType != CB_FORMAT_DATA_RESPONSE || resp.flags != CB_RESPONSE_OK { + t.Fatalf("msgType=%#x flags=%#x", resp.msgType, resp.flags) + } + if v := binary.LittleEndian.Uint32(resp.body); v != DROPEFFECT_COPY { + t.Fatalf("DropEffect=%d, want %d", v, DROPEFFECT_COPY) + } +} + +func TestFormatListAdvertisesDropEffectWithFiles(t *testing.T) { + h, fs := newFileTestHandler(t) + h.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + body := fs.last().body + if !bytes.Contains(body, encodeUTF16LE(FormatNameDropEffect)) { + t.Fatal("staged Format List should advertise Preferred DropEffect") + } +} + +func TestFileContentsRequestServesSizeAndRange(t *testing.T) { + h, fs := newFileTestHandler(t) + data := bytes.Repeat([]byte{0xA5}, 300*1024) // 跨两个 RANGE 段 + h.SetLocalFiles([]LocalFile{{Name: "big.bin", Data: data}}) + + // SIZE 请求 + szReq := make([]byte, 24) + binary.LittleEndian.PutUint32(szReq[0:], 7) // streamId + binary.LittleEndian.PutUint32(szReq[4:], 0) // lindex + binary.LittleEndian.PutUint32(szReq[8:], FILECONTENTS_SIZE) + binary.LittleEndian.PutUint32(szReq[20:], 8) + h.processFileContentsRequest(szReq) + resp := fs.last() + if resp.msgType != CB_FILECONTENTS_RESPONSE || resp.flags != CB_RESPONSE_OK { + t.Fatalf("SIZE: msgType=%#x flags=%#x", resp.msgType, resp.flags) + } + if sid := binary.LittleEndian.Uint32(resp.body); sid != 7 { + t.Fatalf("SIZE: streamId=%d", sid) + } + if size := binary.LittleEndian.Uint64(resp.body[4:]); size != uint64(len(data)) { + t.Fatalf("SIZE: got %d want %d", size, len(data)) + } + + // RANGE 请求(尾段裁剪到文件末尾) + rng := make([]byte, 24) + binary.LittleEndian.PutUint32(rng[0:], 7) + binary.LittleEndian.PutUint32(rng[8:], FILECONTENTS_RANGE) + binary.LittleEndian.PutUint32(rng[12:], 290*1024) // pos + binary.LittleEndian.PutUint32(rng[20:], 1<<20) // 请求远超末尾 + h.processFileContentsRequest(rng) + resp = fs.last() + got := resp.body[4:] + if len(got) != 10*1024 { + t.Fatalf("RANGE: got %d bytes, want clamped %d", len(got), 10*1024) + } + if got[0] != 0xA5 { + t.Fatal("RANGE: wrong payload") + } + + // 越界 lindex → FAIL 且带 streamId + bad := make([]byte, 24) + binary.LittleEndian.PutUint32(bad[0:], 9) + binary.LittleEndian.PutUint32(bad[4:], 5) + binary.LittleEndian.PutUint32(bad[8:], FILECONTENTS_SIZE) + h.processFileContentsRequest(bad) + resp = fs.last() + if resp.flags != CB_RESPONSE_FAIL || len(resp.body) != 4 || + binary.LittleEndian.Uint32(resp.body) != 9 { + t.Fatalf("unknown lindex: flags=%#x body=% X", resp.flags, resp.body) + } +} + +// TestRemoteFileTransferChain 走完 SIZE→RANGE→完成的完整拉取链, +// 并验证服务器裁剪返回段时按实际长度推进。 +func TestRemoteFileTransferChain(t *testing.T) { + h, fs := newFileTestHandler(t) + var resultName string + var resultData []byte + var progGot, progTotal int64 + h.SetFileCallbacks(func(names []string) {}, + func(index int, name string, data []byte) { resultName, resultData = name, data }, + func(index int, received, total int64) { progGot, progTotal = received, total }) + + serverFormatListWithHDrop(t, h) + payload := BuildDropfiles([]string{`E:\doc\说明 书.pdf`}) + h.processFormatDataResponse(payload, CB_RESPONSE_OK) + + file := bytes.Repeat([]byte{0x5A}, fileRangeChunk+1000) + if err := h.RequestRemoteFile(0); err != nil { + t.Fatal(err) + } + + // 1) SIZE 请求 + req := fs.last() + if req.msgType != CB_FILECONTENTS_REQUEST || + binary.LittleEndian.Uint32(req.body[8:]) != FILECONTENTS_SIZE { + t.Fatalf("step1: msgType=%#x flags=%#x", req.msgType, binary.LittleEndian.Uint32(req.body[8:])) + } + sid := binary.LittleEndian.Uint32(req.body[0:]) + // SIZE 响应 → 应发出 RANGE(0, fileRangeChunk) + h.processFileContentsResponse(append(u32le(sid), u64le(uint64(len(file)))...), CB_RESPONSE_OK) + req = fs.last() + if binary.LittleEndian.Uint32(req.body[8:]) != FILECONTENTS_RANGE || + binary.LittleEndian.Uint32(req.body[12:]) != 0 || + binary.LittleEndian.Uint32(req.body[20:]) != fileRangeChunk { + t.Fatalf("step2: want RANGE(0,%d), got pos=%d cb=%d", fileRangeChunk, + binary.LittleEndian.Uint32(req.body[12:]), binary.LittleEndian.Uint32(req.body[20:])) + } + // RANGE 响应(服务器只回一半)→ 应从实际到达位置续传 + h.processFileContentsResponse(append(u32le(sid), file[:fileRangeChunk/2]...), CB_RESPONSE_OK) + req = fs.last() + if p := binary.LittleEndian.Uint32(req.body[12:]); p != fileRangeChunk/2 { + t.Fatalf("step3: want resume pos=%d, got %d", fileRangeChunk/2, p) + } + if progGot != fileRangeChunk/2 || progTotal != int64(len(file)) { + t.Fatalf("progress got=%d total=%d", progGot, progTotal) + } + // 剩余数据 → 完成回调 + h.processFileContentsResponse(append(u32le(sid), file[fileRangeChunk/2:]...), CB_RESPONSE_OK) + if resultName != `E:\doc\说明 书.pdf` { + t.Fatalf("result name=%q", resultName) + } + if !bytes.Equal(resultData, file) { + t.Fatalf("result %d bytes, want %d", len(resultData), len(file)) + } +} + +func TestRemoteFileTransferFailure(t *testing.T) { + h, fs := newFileTestHandler(t) + var failData []byte + var failName string + h.SetFileCallbacks(func(names []string) {}, + func(index int, name string, data []byte) { failName, failData = name, data }, + nil) + serverFormatListWithHDrop(t, h) + h.processFormatDataResponse(BuildDropfiles([]string{"gone.txt"}), CB_RESPONSE_OK) + if err := h.RequestRemoteFile(0); err != nil { + t.Fatal(err) + } + sid := binary.LittleEndian.Uint32(fs.last().body[0:]) + h.processFileContentsResponse(u32le(sid), CB_RESPONSE_FAIL) + if failName != "gone.txt" || failData != nil { + t.Fatalf("failure delivery: name=%q data=%v", failName, failData) + } +} + +// TestServerLocalDropEffectId:Windows 用它本地的 "Preferred DropEffect" +// 注册格式 ID(实测 0xC17E)查询而非我方广告 ID,需按 DropEffect 应答。 +func TestServerLocalDropEffectId(t *testing.T) { + h, fs := newFileTestHandler(t) + h.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + + req := make([]byte, 4) + binary.LittleEndian.PutUint32(req, 0xC17E) + h.processFormatDataRequest(req) + resp := fs.last() + if resp.flags != CB_RESPONSE_OK || len(resp.body) != 4 || + binary.LittleEndian.Uint32(resp.body) != DROPEFFECT_COPY { + t.Fatalf("heuristic DropEffect reply failed: flags=%#x body=% X", resp.flags, resp.body) + } + + // 服务器通告过名称后,精确匹配其本地 ID + h2, fs2 := newFileTestHandler(t) + b := &bytes.Buffer{} + binary.Write(b, binary.LittleEndian, uint32(0xC17E)) + b.Write(encodeUTF16LE(FormatNameDropEffect)) + b.Write([]byte{0, 0}) + h2.processFormatList(b.Bytes(), 0) + h2.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + h2.processFormatDataRequest(req) + if fs2.last().flags != CB_RESPONSE_OK { + t.Fatal("learned serverDropEffectId should be answered with OK") + } + // 无文件时 FAIL + h3, fs3 := newFileTestHandler(t) + b3 := &bytes.Buffer{} + binary.Write(b3, binary.LittleEndian, uint32(0xC17E)) + b3.Write(encodeUTF16LE(FormatNameDropEffect)) + b3.Write([]byte{0, 0}) + h3.processFormatList(b3.Bytes(), 0) + h3.processFormatDataRequest(req) + if fs3.last().flags != CB_RESPONSE_FAIL { + t.Fatal("no staged files: server-local DropEffect should FAIL") + } +} + +func TestRequestRemoteFileOutOfRange(t *testing.T) { + h, _ := newFileTestHandler(t) + if err := h.RequestRemoteFile(3); err == nil { + t.Fatal("expected error for out-of-range index") + } +} + +func TestFormatListAdvertisesFileFormatsWhenStaged(t *testing.T) { + h, fs := newFileTestHandler(t) // long-name 模式 + h.sendFormatList() + body := fs.last().body + if bytes.Contains(body, encodeUTF16LE(FormatNameFileGroupDescriptorW)) { + t.Fatal("file formats must not be advertised before staging files") + } + h.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + body = fs.last().body + if !bytes.Contains(body, encodeUTF16LE(FormatNameFileGroupDescriptorW)) || + !bytes.Contains(body, encodeUTF16LE(FormatNameFileContents)) || + !bytes.Contains(body, encodeUTF16LE(FormatNameDropEffect)) { + t.Fatal("staged Format List should advertise FileGroupDescriptorW/FileContents/DropEffect") + } + // short-name 模式(名称为原始 ASCII 字节) + h2, fs2 := newFileTestHandler(t) + h2.useLongFormatNames = false + h2.SetLocalFiles([]LocalFile{{Name: "a.txt", Data: []byte("hi")}}) + body = fs2.last().body + if !bytes.Contains(body, []byte("FileGroupDescriptorW")) { + t.Fatalf("short-name Format List should carry FileGroupDescriptorW: % X", body[:min(80, len(body))]) + } +} + +func u64le(v uint64) []byte { + b := make([]byte, 8) + binary.LittleEndian.PutUint64(b, v) + return b +} + +// TestParseFileGroupDescriptorWLargeDirectory:远端复制含数千文件的目录 +// 是正常操作(曾因固定 1024 项上限整单被拒);超过防御性上限才拒绝。 +func TestParseFileGroupDescriptorWLargeDirectory(t *testing.T) { + build := func(n uint32) []byte { + body := make([]byte, 4+int(n)*fileDescriptorSize) + binary.LittleEndian.PutUint32(body[0:], n) + for i := uint32(0); i < n; i++ { + d := body[4+int(i)*fileDescriptorSize:] + binary.LittleEndian.PutUint32(d[0:], FD_ATTRIBUTES|FD_FILESIZE|FD_PROGRESSUI) + binary.LittleEndian.PutUint32(d[68:], 7) // 大小低 32 位 + copy(d[72:], asciiUTF16(fmt.Sprintf("f%d.txt", i))) + } + return body + } + // 5000 项:必须接受 + names, _, err := ParseFileGroupDescriptorW(build(5000)) + if err != nil || len(names) != 5000 { + t.Fatalf("5000 items: err=%v names=%d", err, len(names)) + } + if names[4999] != "f4999.txt" { + t.Fatalf("last name=%q", names[4999]) + } + // 超过 maxDescriptorItems:拒绝 + if _, _, err := ParseFileGroupDescriptorW(build(maxDescriptorItems + 1)); err == nil { + t.Fatal("over-limit cItems should be rejected") + } + // cItems 声明与载荷不符:拒绝 + short := build(3)[:4+2*fileDescriptorSize] + binary.LittleEndian.PutUint32(short[0:], 3) + if _, _, err := ParseFileGroupDescriptorW(short); err == nil { + t.Fatal("truncated payload should be rejected") + } +} + +// asciiUTF16 把纯 ASCII 名编码为 UTF-16LE(NUL 结尾),供构造描述符用。 +func asciiUTF16(s string) []byte { + b := make([]byte, 0, (len(s)+1)*2) + for _, c := range []byte(s) { + b = append(b, c, 0) + } + return append(b, 0, 0) +} diff --git a/plugin/cliprdr/handler.go b/plugin/cliprdr/handler.go new file mode 100644 index 0000000..061fb6e --- /dev/null +++ b/plugin/cliprdr/handler.go @@ -0,0 +1,678 @@ +// Package cliprdr handler.go implements a cross-platform CLIPRDR +// (Clipboard Virtual Channel Extension, MS-RDPECLIP) handler for +// bidirectional text clipboard sharing between RDP client and server. +// +// Only text formats (CF_UNICODETEXT / CF_TEXT) are supported. +package cliprdr + +import ( + "bytes" + "encoding/binary" + "log/slog" + "strings" + "sync" + "unicode/utf16" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +// CliprdrHandler implements plugin.ChannelTransport for the "cliprdr" +// static virtual channel. It uses callbacks for clipboard integration +// so that any UI toolkit can wire in its own clipboard access. +type CliprdrHandler struct { + channelSender core.ChannelSender + + useLongFormatNames bool + + // serverCapsReceived is set when the server's CB_CLIP_CAPS PDU has been + // processed. Per MS-RDPECLIP §1.3.2.1 the server sends CB_CLIP_CAPS + // before CB_MONITOR_READY, so FORMAT_LIST should only be sent once both + // have arrived. + serverCapsReceived bool + // monitorReady is set when CB_MONITOR_READY has been received. + monitorReady bool + + // onRemoteClipboardChanged is called with the text when the server's + // clipboard content arrives. + onRemoteClipboardChanged func(text string) + + // getLocalClipboardText is called to retrieve the current local + // clipboard text when the server requests it. + getLocalClipboardText func() string + + // onRemoteClipboardImage is called with PNG-encoded bytes when the + // server's clipboard image (PNG / CF_DIB) arrives. + onRemoteClipboardImage func(png []byte) + + // getLocalClipboardImage is called to retrieve the current local + // clipboard image (PNG-encoded) when the server requests the "PNG" + // registered format. Returns nil when no image is available. + getLocalClipboardImage func() []byte + + // onRemoteClipboardHTML is called with the HTML content when the server's + // clipboard "HTML Format" data arrives (fragment already extracted from + // the CF_HTML envelope). + onRemoteClipboardHTML func(html string) + + // getLocalClipboardHTML is called to retrieve the current local clipboard + // HTML (raw fragment/document, without the CF_HTML envelope) when the + // server requests the "HTML Format" registered format. + // Returns "" when no HTML is available. + getLocalClipboardHTML func() string + + // remoteHTMLFormatId 记录服务器 Format List 中 "HTML Format" 的格式 ID, + // 用于识别 Format Data Response(服务端注册格式的 ID 由服务端分配)。 + remoteHTMLFormatId uint32 + + // serverDropEffectId 记录服务器 Format List 中 "Preferred DropEffect" + // 的服务器本地 ID(注册格式 ID 各端自定,按名称对齐)。 + serverDropEffectId uint32 + + // lastRequestedFormat 记录我方最后一次 Format Data Request 的格式, + // 以正确解读 Format Data Response(PNG 原样、DIB 转码、文本按 UTF-16)。 + lastRequestedFormat uint32 + + // suppressNextLocalChange prevents an echo loop: + // server→client clipboard update triggers a local clipboard change + // event which would otherwise be sent back to the server. + suppressNextLocalChange bool + + // --- File clipboard (CF_HDROP + FileContentsRequest/Response) ----------- + // + // Browser clipboards cannot hold files, so files are "staged" explicitly + // by the UI (file picker / drag-drop): staged files are advertised as + // CF_HDROP in Format List; the server fetches their bytes through + // FileContentsRequest. Files copied on the server arrive as a CF_HDROP + // path list, and the UI pulls bytes per file via RequestRemoteFile. + + mu sync.Mutex + // serverFileClip is set when the server advertised CB_STREAM_FILECLIP_ENABLED. + // CF_HDROP is only advertised/requested when both sides support file streaming. + serverFileClip bool + // localFiles are files staged by the local UI, advertised as CF_HDROP. + localFiles []LocalFile + // remoteFiles are the paths from the server's last file list payload. + remoteFiles []string + // remoteFileSizes are the sizes from the server's FileGroupDescriptorW + // (-1 when unknown; only valid alongside remoteFiles from descriptors). + remoteFileSizes []int64 + // streams tracks in-flight outgoing FileContentsRequest transfers + // (remote → local download) keyed by our StreamId. + streams map[uint32]*fileStream + // nextStreamID is the StreamId counter for outgoing FileContentsRequests. + nextStreamID uint32 + + onRemoteFiles func(names []string) + onFileData func(index int, name string, data []byte) + onFileProgress func(index int, received, total int64) +} + +// NewHandler creates a CliprdrHandler. +// +// - onRemote is called when the server clipboard text is received. +// - getLocal is called to retrieve the current local clipboard text. +// +// Either callback may be nil. Image callbacks can be attached later via +// SetImageCallbacks. +func NewHandler(onRemote func(text string), getLocal func() string) *CliprdrHandler { + return &CliprdrHandler{ + onRemoteClipboardChanged: onRemote, + getLocalClipboardText: getLocal, + streams: map[uint32]*fileStream{}, + } +} + +// SetImageCallbacks attaches image clipboard callbacks (both optional). +// +// - onRemoteImage receives PNG-encoded bytes of the server clipboard image. +// - getLocalImage returns PNG-encoded bytes of the local clipboard image +// (nil when unavailable), used to answer server requests for "PNG". +func (h *CliprdrHandler) SetImageCallbacks(onRemoteImage func(png []byte), getLocalImage func() []byte) { + h.onRemoteClipboardImage = onRemoteImage + h.getLocalClipboardImage = getLocalImage +} + +// SetHTMLCallbacks attaches HTML clipboard callbacks (both optional). +// +// - onRemoteHTML receives the HTML content when the server's clipboard +// "HTML Format" data arrives (fragment extracted from the CF_HTML envelope). +// - getLocalHTML returns the local clipboard HTML (raw HTML, no CF_HTML +// envelope; empty when unavailable), used to answer server requests for +// the "HTML Format" registered format. +func (h *CliprdrHandler) SetHTMLCallbacks(onRemoteHTML func(html string), getLocalHTML func() string) { + h.onRemoteClipboardHTML = onRemoteHTML + h.getLocalClipboardHTML = getLocalHTML +} + +// --- plugin.ChannelTransport interface ------------------------------------ + +func (h *CliprdrHandler) GetType() (string, uint32) { + return ChannelName, ChannelOption +} + +func (h *CliprdrHandler) Sender(f core.ChannelSender) { + h.channelSender = f +} + +// Process handles a reassembled CLIPRDR PDU from the server. +func (h *CliprdrHandler) Process(s []byte) { + if len(s) < 8 { + return + } + r := bytes.NewReader(s) + msgType, _ := core.ReadUint16LE(r) + msgFlags, _ := core.ReadUint16LE(r) + dataLen, _ := core.ReadUInt32LE(r) + + body := make([]byte, dataLen) + if dataLen > 0 { + n, _ := r.Read(body) + body = body[:n] + } + + slog.Debug("cliprdr recv", "msgType", msgType, "msgFlags", msgFlags, "dataLen", dataLen) + + switch msgType { + case CB_CLIP_CAPS: + h.processClipCaps(body) + case CB_MONITOR_READY: + h.processMonitorReady() + case CB_FORMAT_LIST: + h.processFormatList(body, msgFlags) + case CB_FORMAT_LIST_RESPONSE: + h.processFormatListResponse(msgFlags) + case CB_FORMAT_DATA_REQUEST: + h.processFormatDataRequest(body) + case CB_FORMAT_DATA_RESPONSE: + h.processFormatDataResponse(body, msgFlags) + case CB_FILECONTENTS_REQUEST: + h.processFileContentsRequest(body) + case CB_FILECONTENTS_RESPONSE: + h.processFileContentsResponse(body, msgFlags) + case CB_LOCK_CLIPDATA, CB_UNLOCK_CLIPDATA: + // ignored + default: + slog.Debug("cliprdr: unhandled msgType", "msgType", msgType) + } +} + +// --- Clipboard Capabilities (MS-RDPECLIP 2.2.2.1) ------------------------- + +func (h *CliprdrHandler) processClipCaps(body []byte) { + if len(body) < 4 { + return + } + cCapSets := binary.LittleEndian.Uint16(body[0:2]) + // pad1 at [2:4] + offset := 4 + for i := 0; i < int(cCapSets); i++ { + if offset+4 > len(body) { + break + } + capType := binary.LittleEndian.Uint16(body[offset:]) + capLen := binary.LittleEndian.Uint16(body[offset+2:]) + if capType == CB_CAPSTYPE_GENERAL && capLen >= 12 { + generalFlags := binary.LittleEndian.Uint32(body[offset+8:]) + h.useLongFormatNames = generalFlags&CB_USE_LONG_FORMAT_NAMES != 0 + h.mu.Lock() + h.serverFileClip = generalFlags&CB_STREAM_FILECLIP_ENABLED != 0 + fileClip := h.serverFileClip + h.mu.Unlock() + slog.Debug("cliprdr: server caps", "generalFlags", generalFlags, + "longNames", h.useLongFormatNames, "fileClip", fileClip) + } + offset += int(capLen) + } + h.serverCapsReceived = true + // If CB_MONITOR_READY already arrived before CB_CLIP_CAPS (non-standard + // ordering), send the FORMAT_LIST now that we have correct capabilities. + if h.monitorReady { + h.sendFormatList() + } +} + +func (h *CliprdrHandler) sendClipCaps() { + b := &bytes.Buffer{} + // General capability set: type(2) + length(2) + version(4) + flags(4) + binary.Write(b, binary.LittleEndian, uint16(CB_CAPSTYPE_GENERAL)) + binary.Write(b, binary.LittleEndian, uint16(12)) + binary.Write(b, binary.LittleEndian, uint32(CB_CAPS_VERSION_2)) + // CB_STREAM_FILECLIP_ENABLED + CB_FILECLIP_NO_FILE_PATHS:声明支持 + // FileGroupDescriptorW/FileContents 流式文件传输(mstsc 同款), + // 服务器未开时其会忽略这些位。 + binary.Write(b, binary.LittleEndian, + uint32(CB_USE_LONG_FORMAT_NAMES|CB_STREAM_FILECLIP_ENABLED|CB_FILECLIP_NO_FILE_PATHS)) + + // cCapabilitySets(2) + pad1(2) + capabilitySet + body := &bytes.Buffer{} + binary.Write(body, binary.LittleEndian, uint16(1)) + binary.Write(body, binary.LittleEndian, uint16(0)) + body.Write(b.Bytes()) + + h.sendPDU(CB_CLIP_CAPS, 0, body.Bytes()) +} + +// --- Monitor Ready (MS-RDPECLIP 2.2.2.2) ---------------------------------- + +func (h *CliprdrHandler) processMonitorReady() { + slog.Debug("cliprdr: server Monitor Ready") + h.monitorReady = true + h.sendClipCaps() + // Per MS-RDPECLIP §1.3.2.1 the server sends CB_CLIP_CAPS before + // CB_MONITOR_READY. Only send FORMAT_LIST after server caps are known + // so useLongFormatNames is set correctly. If CB_CLIP_CAPS hasn't been + // received yet (non-standard ordering), defer until processClipCaps fires. + if h.serverCapsReceived { + h.sendFormatList() + } +} + +// --- Format List (MS-RDPECLIP 2.2.3.1) ------------------------------------ + +func (h *CliprdrHandler) sendFormatList() { + h.mu.Lock() + staged := len(h.localFiles) > 0 + fileClip := h.serverFileClip + h.mu.Unlock() + // 仅当双方都声明 CB_STREAM_FILECLIP_ENABLED 且本地已暂存文件时 + // 才广告 CF_HDROP,避免服务器把普通粘贴当文件处理。 + advertiseFiles := staged && fileClip + + b := &bytes.Buffer{} + if h.useLongFormatNames { + // Long Format Name: formatId(4) + wszFormatName(null-terminated UTF-16LE) + writeFormat := func(id uint32, name string) { + binary.Write(b, binary.LittleEndian, id) + b.Write(encodeUTF16LE(name)) + b.Write([]byte{0, 0}) // null terminator + } + writeFormat(CF_UNICODETEXT, "") // 空名 = 标准格式 + writeFormat(CF_PNG, FormatNamePNG) // 注册格式 "PNG" + writeFormat(CF_HTML_FORMAT_ID, FormatNameHTML) // 注册格式 "HTML Format" + if advertiseFiles { + // 文件传输三件套(mstsc 同款,名称必须一致): + // 描述符列表 + 流式内容 + 复制语义 + writeFormat(CF_FILE_GROUP_DESCRIPTORW, FormatNameFileGroupDescriptorW) + writeFormat(CF_FILE_CONTENTS, FormatNameFileContents) + writeFormat(CF_DROP_EFFECT, FormatNameDropEffect) + } + } else { + // Short Format Name: formatId(4) + formatName[32] + writeShort := func(id uint32, name string) { + binary.Write(b, binary.LittleEndian, id) + nameBuf := make([]byte, 32) + copy(nameBuf, name) + b.Write(nameBuf) + } + writeShort(CF_UNICODETEXT, "") + writeShort(CF_PNG, FormatNamePNG) + writeShort(CF_HTML_FORMAT_ID, FormatNameHTML) + if advertiseFiles { + writeShort(CF_FILE_GROUP_DESCRIPTORW, FormatNameFileGroupDescriptorW) + writeShort(CF_FILE_CONTENTS, FormatNameFileContents) + writeShort(CF_DROP_EFFECT, FormatNameDropEffect) + } + } + h.sendPDU(CB_FORMAT_LIST, 0, b.Bytes()) +} + +func (h *CliprdrHandler) processFormatList(body []byte, msgFlags uint16) { + formats := h.parseFormatList(body, msgFlags) + slog.Debug("cliprdr: server Format List", "formats", formats) + + // 记录服务器侧 "Preferred DropEffect" 的本地 ID(注册格式 ID 各端自定) + for _, f := range formats { + if strings.EqualFold(f.FormatName, FormatNameDropEffect) { + h.mu.Lock() + h.serverDropEffectId = f.FormatId + h.mu.Unlock() + } + } + + // Always respond OK + h.sendPDU(CB_FORMAT_LIST_RESPONSE, CB_RESPONSE_OK, nil) + + // 请求优先级:文件描述符(远端复制了文件)> CF_HDROP > PNG > CF_DIB > + // HTML > 文本。资源管理器复制文件时通常还附带文件名文本,文件列表 + // 必须最先识别。 + for _, f := range formats { + if strings.EqualFold(f.FormatName, FormatNameFileGroupDescriptorW) { + h.mu.Lock() + fileClip := h.serverFileClip + h.mu.Unlock() + if fileClip { + slog.Debug("cliprdr: server offers files (FileGroupDescriptorW)") + h.lastRequestedFormat = CF_FILE_GROUP_DESCRIPTORW + h.sendFormatDataRequest(f.FormatId) + return + } + } + } + for _, f := range formats { + if f.FormatId == uint32(CF_HDROP) || strings.EqualFold(f.FormatName, FormatNameHDrop) { + h.mu.Lock() + fileClip := h.serverFileClip + h.mu.Unlock() + if fileClip { + slog.Debug("cliprdr: server offers files (CF_HDROP)") + h.lastRequestedFormat = CF_HDROP + h.sendFormatDataRequest(CF_HDROP) + return + } + } + } + for _, f := range formats { + if f.FormatId == CF_PNG || strings.EqualFold(f.FormatName, FormatNamePNG) || + strings.EqualFold(f.FormatName, "image/png") { + h.lastRequestedFormat = f.FormatId + h.sendFormatDataRequest(f.FormatId) + return + } + } + for _, f := range formats { + if f.FormatId == CF_DIB { + h.lastRequestedFormat = CF_DIB + h.sendFormatDataRequest(CF_DIB) + return + } + } + // HTML Format:远端复制的富文本(浏览器/Office),保格式优于纯文本 + for _, f := range formats { + if f.FormatId == CF_HTML_FORMAT_ID || strings.EqualFold(f.FormatName, FormatNameHTML) { + h.remoteHTMLFormatId = f.FormatId + h.lastRequestedFormat = f.FormatId + h.sendFormatDataRequest(f.FormatId) + return + } + } + for _, f := range formats { + if f.FormatId == CF_UNICODETEXT { + h.lastRequestedFormat = CF_UNICODETEXT + h.sendFormatDataRequest(CF_UNICODETEXT) + return + } + } + for _, f := range formats { + if f.FormatId == CF_TEXT { + h.lastRequestedFormat = CF_TEXT + h.sendFormatDataRequest(CF_TEXT) + return + } + } +} + +func (h *CliprdrHandler) parseFormatList(body []byte, msgFlags uint16) []CliprdrFormat { + var formats []CliprdrFormat + if h.useLongFormatNames && (msgFlags&CB_ASCII_NAMES == 0) { + // Long Format Names (MS-RDPECLIP 2.2.3.1.1.1) + offset := 0 + for offset+4 <= len(body) { + fmtId := binary.LittleEndian.Uint32(body[offset:]) + offset += 4 + // Read null-terminated UTF-16LE string + nameEnd := offset + for nameEnd+1 < len(body) { + if body[nameEnd] == 0 && body[nameEnd+1] == 0 { + break + } + nameEnd += 2 + } + name := decodeUTF16LE(body[offset:nameEnd]) + offset = nameEnd + 2 + formats = append(formats, CliprdrFormat{fmtId, name}) + } + } else { + // Short Format Names (MS-RDPECLIP 2.2.3.1.1.2) + offset := 0 + for offset+36 <= len(body) { + fmtId := binary.LittleEndian.Uint32(body[offset:]) + nameBytes := body[offset+4 : offset+36] + var name string + if msgFlags&CB_ASCII_NAMES != 0 { + name = strings.TrimRight(string(nameBytes), "\x00") + } else { + name = decodeUTF16LE(nameBytes) + name = strings.TrimRight(name, "\x00") + } + formats = append(formats, CliprdrFormat{fmtId, name}) + offset += 36 + } + } + return formats +} + +func (h *CliprdrHandler) processFormatListResponse(msgFlags uint16) { + if msgFlags&CB_RESPONSE_OK != 0 { + slog.Debug("cliprdr: Format List Response OK") + } else { + slog.Warn("cliprdr: Format List Response FAIL") + } +} + +// --- Format Data Request / Response (MS-RDPECLIP 2.2.5) -------------------- + +func (h *CliprdrHandler) sendFormatDataRequest(formatId uint32) { + b := make([]byte, 4) + binary.LittleEndian.PutUint32(b, formatId) + h.sendPDU(CB_FORMAT_DATA_REQUEST, 0, b) + slog.Debug("cliprdr: sent Format Data Request", "formatId", formatId) +} + +func (h *CliprdrHandler) processFormatDataRequest(body []byte) { + if len(body) < 4 { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + requestedFormat := binary.LittleEndian.Uint32(body[0:4]) + slog.Debug("cliprdr: server requests format", "formatId", requestedFormat) + + switch requestedFormat { + case CF_PNG: + var png []byte + if h.getLocalClipboardImage != nil { + png = h.getLocalClipboardImage() + } + if len(png) == 0 { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, png) + case CF_HTML_FORMAT_ID: + // 本地 HTML(无 CF_HTML 信封)→ 包装成 "HTML Format" 字节流 + html := "" + if h.getLocalClipboardHTML != nil { + html = h.getLocalClipboardHTML() + } + if html == "" { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, EncodeCFHTML(html)) + case CF_FILE_GROUP_DESCRIPTORW: + // 远端粘贴文件:回 FILE_GROUP_DESCRIPTORW(名称+大小+属性), + // 字节内容由服务器随后以 CB_FILECONTENTS_REQUEST 按 lIndex 拉取 + h.mu.Lock() + files := h.localFiles + h.mu.Unlock() + if len(files) == 0 { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, BuildFileGroupDescriptorW(files)) + case CF_FILE_CONTENTS: + // FileContents 的具体字节一律走 CB_FILECONTENTS_REQUEST(含流式 + // 偏移/长度),对整格式 Format Data Request 只能回答不可流式的空体 + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + case CF_DROP_EFFECT: + // 有暂存文件时回答 DROPEFFECT_COPY(否则 FAIL)。 + // 该查询决定资源管理器右键菜单"粘贴"是否可用。 + h.mu.Lock() + hasFiles := len(h.localFiles) > 0 + h.mu.Unlock() + if !hasFiles { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, buildDropEffectReply()) + case CF_UNICODETEXT: + text := "" + if h.getLocalClipboardText != nil { + text = h.getLocalClipboardText() + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, encodeUTF16LE(text+"\x00")) + case CF_TEXT: + text := "" + if h.getLocalClipboardText != nil { + text = h.getLocalClipboardText() + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, []byte(text+"\x00")) + default: + // 注册格式 ID 各端自定:服务器会用它本地的 "Preferred DropEffect" + // ID 查询(不一定先通告)。该查询决定资源管理器"粘贴"是否可用。 + h.mu.Lock() + serverDropEffect := h.serverDropEffectId + hasFiles := len(h.localFiles) > 0 + h.mu.Unlock() + + // DropEffect 应答:精确命中服务器 ID,或未知注册格式 ID(≥0xC000) + // 且有暂存文件时按 DropEffect 猜测应答 + if (serverDropEffect != 0 && requestedFormat == serverDropEffect) || + (requestedFormat >= 0xC000 && hasFiles) { + if !hasFiles { + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + return + } + slog.Debug("cliprdr: replying DropEffect", "formatId", requestedFormat) + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_OK, buildDropEffectReply()) + return + } + h.sendPDU(CB_FORMAT_DATA_RESPONSE, CB_RESPONSE_FAIL, nil) + } +} + +func (h *CliprdrHandler) processFormatDataResponse(body []byte, msgFlags uint16) { + if msgFlags&CB_RESPONSE_OK == 0 { + slog.Warn("cliprdr: Format Data Response FAIL") + return + } + + // 按我方请求的格式解读响应 + switch { + case h.lastRequestedFormat == CF_FILE_GROUP_DESCRIPTORW: + names, sizes, err := ParseFileGroupDescriptorW(body) + if err != nil { + slog.Warn("cliprdr: FileGroupDescriptorW parse failed", "err", err, "len", len(body)) + return + } + h.mu.Lock() + h.remoteFiles = names + h.remoteFileSizes = sizes + h.mu.Unlock() + if h.onRemoteFiles != nil { + h.onRemoteFiles(names) + } + case h.lastRequestedFormat == CF_HDROP: + names, err := ParseDropfiles(body) + if err != nil { + slog.Warn("cliprdr: CF_HDROP parse failed", "err", err, "len", len(body)) + return + } + h.mu.Lock() + h.remoteFiles = names + h.mu.Unlock() + if h.onRemoteFiles != nil { + h.onRemoteFiles(names) + } + case h.lastRequestedFormat != 0 && h.lastRequestedFormat == h.remoteHTMLFormatId: + // "HTML Format"(服务端注册格式的 ID 与请求时记录的一致) + html := DecodeCFHTML(body) + if html == "" { + return + } + if h.onRemoteClipboardHTML != nil { + h.suppressNextLocalChange = true + h.onRemoteClipboardHTML(html) + } + case h.lastRequestedFormat == CF_PNG: + if len(body) == 0 { + return + } + if h.onRemoteClipboardImage != nil { + h.suppressNextLocalChange = true + png := make([]byte, len(body)) + copy(png, body) + h.onRemoteClipboardImage(png) + } + case h.lastRequestedFormat == CF_DIB || h.lastRequestedFormat == CF_DIBV5: + png, err := DIBToPNG(body) + if err != nil { + slog.Warn("cliprdr: DIB→PNG conversion failed", "err", err) + return + } + if h.onRemoteClipboardImage != nil { + h.suppressNextLocalChange = true + h.onRemoteClipboardImage(png) + } + default: + // CF_UNICODETEXT / CF_TEXT → UTF-16LE 文本 + text := decodeUTF16LE(body) + text = strings.TrimRight(text, "\x00") + if text != "" && h.onRemoteClipboardChanged != nil { + slog.Debug("cliprdr: received text", "len", len(text)) + h.suppressNextLocalChange = true + h.onRemoteClipboardChanged(text) + } + } +} + +// --- Public API for local clipboard changes -------------------------------- + +// OnLocalClipboardChanged notifies the server that the local clipboard +// content has changed. Call this from the UI when the system clipboard +// changes (e.g. via polling or a platform clipboard-change signal). +func (h *CliprdrHandler) OnLocalClipboardChanged() { + if h.suppressNextLocalChange { + h.suppressNextLocalChange = false + return + } + if h.channelSender != nil { + h.sendFormatList() + slog.Debug("cliprdr: local clipboard changed, sent Format List") + } +} + +// --- Send helpers ---------------------------------------------------------- + +func (h *CliprdrHandler) sendPDU(msgType, msgFlags uint16, body []byte) { + if h.channelSender == nil { + return + } + sendClipPDU(h.channelSender, msgType, msgFlags, body) +} + +// --- UTF-16LE helpers ------------------------------------------------------ + +func decodeUTF16LE(b []byte) string { + if len(b) < 2 { + return "" + } + // Trim to even length + if len(b)%2 != 0 { + b = b[:len(b)-1] + } + u16 := make([]uint16, len(b)/2) + for i := range u16 { + u16[i] = binary.LittleEndian.Uint16(b[i*2:]) + } + return string(utf16.Decode(u16)) +} + +func encodeUTF16LE(s string) []byte { + runes := []rune(s) + u16 := utf16.Encode(runes) + b := make([]byte, len(u16)*2) + for i, v := range u16 { + binary.LittleEndian.PutUint16(b[i*2:], v) + } + return b +} diff --git a/plugin/cliprdr/html_format.go b/plugin/cliprdr/html_format.go new file mode 100644 index 0000000..ad8768c --- /dev/null +++ b/plugin/cliprdr/html_format.go @@ -0,0 +1,77 @@ +package cliprdr + +import ( + "fmt" + "strings" +) + +// Windows "HTML Format" 剪贴板编解码(MS-Doc: HTML Clipboard Format)。 +// +// 字节流 = 固定文本头 + HTML 文档。头共 5 行、每行 \r\n 结尾: +// +// Version:0.9\r\n (13 字节) +// StartHTML:0000000105\r\n (22 字节) +// EndHTML:0000000256\r\n (20 字节) +// StartFragment:0000000139\r\n (26 字节) +// EndFragment:0000000222\r\n (24 字节) +// +// 偏移为 10 位零填充的十进制**字节**偏移(UTF-8)。固定头总长 105 字节。 + +const cfHTMLHeaderLen = 13 + 22 + 20 + 26 + 24 // 105 + +const cfHTMLWrapPrefix = "" +const cfHTMLWrapSuffix = "" + +// EncodeCFHTML 把 HTML 片段(fragment)包装成 "HTML Format" 剪贴板字节流。 +func EncodeCFHTML(fragment string) []byte { + startHTML := cfHTMLHeaderLen + startFragment := startHTML + len(cfHTMLWrapPrefix) + endFragment := startFragment + len(fragment) + endHTML := endFragment + len(cfHTMLWrapSuffix) + + head := fmt.Sprintf("Version:0.9\r\nStartHTML:%010d\r\nEndHTML:%010d\r\nStartFragment:%010d\r\nEndFragment:%010d\r\n", + startHTML, endHTML, startFragment, endFragment) + doc := cfHTMLWrapPrefix + fragment + cfHTMLWrapSuffix + return []byte(head + doc) +} + +// DecodeCFHTML 从 "HTML Format" 剪贴板字节流中提取 StartFragment 到 +// EndFragment 之间的 HTML 内容。头部解析失败时回退为整段字符串。 +func DecodeCFHTML(data []byte) string { + s := string(data) + lines := strings.Split(s, "\r\n") + var startFragment, endFragment int = -1, -1 + // 头最多 5 行;解析到两个 Fragment 偏移即可停止。 + for i, ln := range lines { + if i >= 5 { + break + } + if v, ok := parseOffsetLine(ln, "StartFragment:"); ok { + startFragment = v + } else if v, ok := parseOffsetLine(ln, "EndFragment:"); ok { + endFragment = v + } + if startFragment >= 0 && endFragment > startFragment && endFragment <= len(data) { + break + } + } + if startFragment < 0 || endFragment <= startFragment || endFragment > len(data) { + // 回退:无有效偏移时原样返回(去除首行头部不可行,尽力而为) + return s + } + return s[startFragment:endFragment] +} + +func parseOffsetLine(line, key string) (int, bool) { + if !strings.HasPrefix(line, key) { + return 0, false + } + v := 0 + for _, c := range line[len(key):] { + if c < '0' || c > '9' { + return 0, false + } + v = v*10 + int(c-'0') + } + return v, true +} diff --git a/plugin/cliprdr/html_format_test.go b/plugin/cliprdr/html_format_test.go new file mode 100644 index 0000000..04e793f --- /dev/null +++ b/plugin/cliprdr/html_format_test.go @@ -0,0 +1,104 @@ +package cliprdr + +import ( + "bytes" + "encoding/binary" + "strings" + "testing" +) + +func TestEncodeCFHTMLRoundTrip(t *testing.T) { + frag := "

红色 加粗

" + data := EncodeCFHTML(frag) + + s := string(data) + if !strings.HasPrefix(s, "Version:0.9\r\n") { + t.Fatalf("missing version header: %q", s[:30]) + } + // 固定头长 105 字节,doc 从 105 开始 + if !strings.HasPrefix(s[105:], "") { + t.Fatalf("html doc should start at offset 105 with wrapper prefix") + } + // 偏移字段自洽(10 位零填充、字节偏移) + expect := func(key string, want int) { + i := strings.Index(s, key) + if i < 0 { + t.Fatalf("missing %s", key) + } + v := 0 + for _, c := range s[i+len(key) : i+len(key)+10] { + if c < '0' || c > '9' { + t.Fatalf("%s offset not 10-digit: %q", key, s[i+len(key):i+len(key)+10]) + } + v = v*10 + int(c-'0') + } + if v != want { + t.Fatalf("%s=%d, want %d", key, v, want) + } + } + expect("StartHTML:", 105) + expect("EndHTML:", len(data)) + expect("StartFragment:", 105+len(cfHTMLWrapPrefix)) + expect("EndFragment:", len(data)-len(cfHTMLWrapSuffix)) + + // 往返 + if got := DecodeCFHTML(data); got != frag { + t.Fatalf("round-trip mismatch:\n got %q\nwant %q", got, frag) + } + // 二进制读取不消费编码流(防御:编码流是文本格式) + var n uint32 + if err := binary.Read(bytes.NewReader(data[:4]), binary.LittleEndian, &n); err != nil { + t.Fatalf("binary sanity: %v", err) + } +} + +func TestDecodeCFHTMLWindowsSample(t *testing.T) { + // 模拟 Windows 剪贴板的实际字节流(头部 + 文档),含 UTF-8 中文 + frag := "

hello \xe4\xb8\xad\xe6\x96\x87

" + doc := "" + frag + "" + var buf bytes.Buffer + buf.WriteString("Version:0.9\r\nStartHTML:0000000105\r\n") + endHTML := 105 + len(doc) + buf.WriteString("EndHTML:" + pad10(endHTML) + "\r\n") + startFragment := 105 + len(cfHTMLWrapPrefix) + buf.WriteString("StartFragment:" + pad10(startFragment) + "\r\n") + endFragment := startFragment + len(frag) + buf.WriteString("EndFragment:" + pad10(endFragment) + "\r\n") + buf.WriteString(doc) + + got := DecodeCFHTML(buf.Bytes()) + if got != frag { + t.Fatalf("decode mismatch:\n got %q\nwant %q", got, frag) + } +} + +func TestDecodeCFHTMLFallback(t *testing.T) { + // 无有效头的乱数据:回退为原样返回,不得 panic + got := DecodeCFHTML([]byte("not a cf_html stream")) + if got != "not a cf_html stream" { + t.Fatalf("fallback changed content: %q", got) + } +} + +func pad10(v int) string { + s := itoa10(v) + if len(s) > 10 { + s = s[len(s)-10:] + } + for len(s) < 10 { + s = "0" + s + } + return s +} + +func itoa10(v int) string { + if v == 0 { + return "0" + } + var b []byte + for v > 0 { + b = append([]byte{byte('0' + v%10)}, b...) + v /= 10 + } + return string(b) +} diff --git a/plugin/cliprdr/image_dib.go b/plugin/cliprdr/image_dib.go new file mode 100644 index 0000000..43e0c0e --- /dev/null +++ b/plugin/cliprdr/image_dib.go @@ -0,0 +1,114 @@ +// image_dib.go 将 Windows 剪贴板的 CF_DIB(BITMAPINFO + 像素)转换为 PNG。 +// +// 剪贴板截图类内容常见两种来源:现代应用直接提供 "PNG" 注册格式;老应用 +// (如 Server 上的画图)只提供 CF_DIB。本转换器覆盖最常见子集: +// BI_RGB / BI_BITFIELDS 压缩、24/32bpp、底行优先(默认)与 top-down。 +package cliprdr + +import ( + "bytes" + "encoding/binary" + "fmt" + "image" + "image/png" +) + +// DIBToPNG 把 CF_DIB 数据解码并编码为 PNG;不支持的子集返回错误。 +func DIBToPNG(dib []byte) ([]byte, error) { + if len(dib) < 40 { + return nil, fmt.Errorf("dib too short: %d", len(dib)) + } + biSize := binary.LittleEndian.Uint32(dib[0:4]) + w := int(int32(binary.LittleEndian.Uint32(dib[4:8]))) + hRaw := int32(binary.LittleEndian.Uint32(dib[8:12])) + topDown := hRaw < 0 + h := int(hRaw) + if h < 0 { + h = -h + } + bpp := binary.LittleEndian.Uint16(dib[14:16]) + compression := binary.LittleEndian.Uint32(dib[16:20]) + if w <= 0 || h <= 0 || w > 16384 || h > 16384 { + return nil, fmt.Errorf("bad dimensions %dx%d", w, h) + } + + var rMask, gMask, bMask uint32 + pixOff := int(biSize) + switch compression { + case 0: // BI_RGB + rMask, gMask, bMask = 0x00FF0000, 0x0000FF00, 0x000000FF + case 3: // BI_BITFIELDS + if biSize >= 108 { // BITMAPV4HEADER/V5HEADER 内含掩码 + rMask = binary.LittleEndian.Uint32(dib[40:44]) + gMask = binary.LittleEndian.Uint32(dib[44:48]) + bMask = binary.LittleEndian.Uint32(dib[48:52]) + } else { // 三个 DWORD 掩码紧跟标准头 + if pixOff+12 > len(dib) { + return nil, fmt.Errorf("missing bitfields masks") + } + rMask = binary.LittleEndian.Uint32(dib[pixOff:]) + gMask = binary.LittleEndian.Uint32(dib[pixOff+4:]) + bMask = binary.LittleEndian.Uint32(dib[pixOff+8:]) + pixOff += 12 + } + default: + return nil, fmt.Errorf("unsupported compression %d", compression) + } + if bpp != 24 && bpp != 32 { + return nil, fmt.Errorf("unsupported bpp %d", bpp) + } + + stride := (w*int(bpp) + 31) / 32 * 4 + if pixOff+stride*h > len(dib) { + return nil, fmt.Errorf("dib truncated: need %d, have %d", pixOff+stride*h, len(dib)) + } + + img := image.NewRGBA(image.Rect(0, 0, w, h)) + for y := 0; y < h; y++ { + srcY := y + if !topDown { + srcY = h - 1 - y + } + row := dib[pixOff+srcY*stride:] + for x := 0; x < w; x++ { + var r8, g8, b8 uint8 + if bpp == 32 { + v := binary.LittleEndian.Uint32(row[x*4:]) + r8, g8, b8 = maskTo8(v, rMask), maskTo8(v, gMask), maskTo8(v, bMask) + } else { + o := x * 3 + b8, g8, r8 = row[o], row[o+1], row[o+2] + } + i := img.PixOffset(x, y) + img.Pix[i], img.Pix[i+1], img.Pix[i+2], img.Pix[i+3] = r8, g8, b8, 255 + } + } + + var buf bytes.Buffer + if err := png.Encode(&buf, img); err != nil { + return nil, err + } + return buf.Bytes(), nil +} + +// maskTo8 把 mask 提取出的位域缩放为 8 位 +func maskTo8(v, mask uint32) uint8 { + if mask == 0 { + return 0 + } + shift := 0 + for mask&1 == 0 { + mask >>= 1 + shift++ + } + bits := 0 + for mask != 0 { + mask >>= 1 + bits++ + } + val := (v >> uint(shift)) & ((1 << uint(bits)) - 1) + if bits < 8 { + return uint8(val<<(8-bits) | val>>(2*uint(bits)-8)) + } + return uint8(val >> uint(bits-8)) +} diff --git a/plugin/drdynvc/dvc.go b/plugin/drdynvc/dvc.go new file mode 100644 index 0000000..6c4cc0e --- /dev/null +++ b/plugin/drdynvc/dvc.go @@ -0,0 +1,364 @@ +package drdynvc + +import ( + "bytes" + "encoding/hex" + "io" + "log/slog" + "strings" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/plugin" +) + +const ( + ChannelName = plugin.DRDYNVC_SVC_CHANNEL_NAME + ChannelOption = plugin.CHANNEL_OPTION_INITIALIZED | + plugin.CHANNEL_OPTION_ENCRYPT_RDP +) + +const ( + MAX_DVC_CHANNELS = 20 +) + +const ( + DYNVC_CREATE_REQ = 0x01 + DYNVC_DATA_FIRST = 0x02 + DYNVC_DATA = 0x03 + DYNVC_CLOSE = 0x04 + DYNVC_CAPABILITIES = 0x05 + DYNVC_DATA_FIRST_COMPRESSED = 0x06 + DYNVC_DATA_COMPRESSED = 0x07 + DYNVC_SOFT_SYNC_REQUEST = 0x08 + DYNVC_SOFT_SYNC_RESPONSE = 0x09 +) + +// DvcChannelHandler processes data for a specific dynamic virtual channel. +type DvcChannelHandler interface { + Process(data []byte) +} + +type ChannelClient struct { + name string + id uint32 + channelSender core.ChannelSender +} + +type dvcChannelInfo struct { + name string + id uint32 + cbChId uint8 + handler DvcChannelHandler +} + +type dvcReassembly struct { + buf bytes.Buffer + totalLen uint32 +} + +type DvcClient struct { + w core.ChannelSender + channels map[string]ChannelClient + handlers map[string]DvcChannelHandler // channelName → handler + rejectedChannels map[string]bool // channelName → explicitly rejected + channelById map[uint32]*dvcChannelInfo // channelId → info + reassembly map[uint32]*dvcReassembly // channelId → reassembly state + negotiatedVersion uint16 +} + +func NewDvcClient() *DvcClient { + return &DvcClient{ + channels: make(map[string]ChannelClient, 100), + handlers: make(map[string]DvcChannelHandler), + rejectedChannels: make(map[string]bool), + channelById: make(map[uint32]*dvcChannelInfo), + reassembly: make(map[uint32]*dvcReassembly), + } +} + +// RegisterHandler registers a handler for a named DVC channel. +func (c *DvcClient) RegisterHandler(name string, handler DvcChannelHandler) { + c.handlers[name] = handler +} + +// RegisterRejectedChannel marks a DVC channel to be explicitly rejected +// (non-zero CreationStatus) so the server does not use it. +// Use this to steer servers toward a fallback channel; for example, +// rejecting AUDIO_PLAYBACK_LOSSY_DVC forces gnome-remote-desktop to +// fall back to lossless AUDIO_PLAYBACK_DVC (PCM). +func (c *DvcClient) RegisterRejectedChannel(name string) { + c.rejectedChannels[name] = true +} + +func (c *DvcClient) LoadAddin(f core.ChannelSender) { + +} + +type DvcHeader struct { + cmd uint8 + sp uint8 + cbChId uint8 +} + +func readHeader(r io.Reader) *DvcHeader { + value, _ := core.ReadUInt8(r) + cmd := (value & 0xf0) >> 4 + sp := (value & 0x0c) >> 2 + cbChId := (value & 0x03) >> 0 + return &DvcHeader{cmd, sp, cbChId} +} + +func (h *DvcHeader) serialize(channelId uint32) []byte { + b := &bytes.Buffer{} + core.WriteUInt8((h.cmd<<4)|(h.sp<<2)|h.cbChId, b) + if h.cbChId == 0 { + core.WriteUInt8(uint8(channelId), b) + } else if h.cbChId == 1 { + core.WriteUInt16LE(uint16(channelId), b) + } else { + core.WriteUInt32LE(channelId, b) + } + + return b.Bytes() +} + +func (c *DvcClient) Send(s []byte) (int, error) { + slog.Debug("dvc Send", "len", len(s), "data", hex.EncodeToString(s)) + name, _ := c.GetType() + return c.w.SendToChannel(name, s) +} + +// SendDvcData sends data on a DVC channel wrapped in a DYNVC_DATA PDU. +func (c *DvcClient) SendDvcData(channelId uint32, data []byte) { + ch, ok := c.channelById[channelId] + if !ok { + return + } + hdr := &DvcHeader{cmd: DYNVC_DATA, sp: 0, cbChId: ch.cbChId} + b := &bytes.Buffer{} + b.Write(hdr.serialize(channelId)) + b.Write(data) + c.Send(b.Bytes()) +} +func (c *DvcClient) Sender(f core.ChannelSender) { + c.w = f +} +func (c *DvcClient) GetType() (string, uint32) { + return ChannelName, ChannelOption +} + +func (c *DvcClient) Process(s []byte) { + defer func() { + if r := recover(); r != nil { + slog.Error("dvc: panic in Process", "err", r) + } + }() + r := bytes.NewReader(s) + hdr := readHeader(r) + b, _ := core.ReadBytes(r.Len(), r) + + switch hdr.cmd { + case DYNVC_CAPABILITIES: + slog.Debug("DYNVC_CAPABILITIES") + c.processCapsPdu(hdr, b) + case DYNVC_CREATE_REQ: + slog.Debug("DYNVC_CREATE_REQ") + c.processCreateReq(hdr, b) + case DYNVC_DATA_FIRST: + c.processDataFirst(hdr, b) + case DYNVC_DATA: + c.processData(hdr, b) + case DYNVC_CLOSE: + c.processClose(hdr, b) + case DYNVC_SOFT_SYNC_REQUEST: + slog.Debug("DYNVC_SOFT_SYNC_REQUEST") + c.processSoftSyncRequest(hdr, b) + default: + slog.Warn("dvc: unhandled cmd", "cmd", hdr.cmd) + } +} +func (c *DvcClient) processClose(hdr *DvcHeader, s []byte) { + r := bytes.NewReader(s) + channelId := readDvcId(r, hdr.cbChId) + ch, ok := c.channelById[channelId] + name := "(unknown)" + if ok { + name = ch.name + delete(c.channelById, channelId) + delete(c.reassembly, channelId) + } + slog.Debug("dvc: CLOSE", "channelId", channelId, "name", name) +} + +func (c *DvcClient) processCreateReq(hdr *DvcHeader, s []byte) { + r := bytes.NewReader(s) + channelId := readDvcId(r, hdr.cbChId) + nameBytes, _ := core.ReadBytes(r.Len(), r) + channelName := strings.TrimRight(string(nameBytes), "\x00") + slog.Debug("dvc: create request", "channelId", channelId, "name", channelName) + + // Associate handler if registered + var handler DvcChannelHandler + if h, ok := c.handlers[channelName]; ok { + handler = h + info := &dvcChannelInfo{ + name: channelName, + id: channelId, + cbChId: hdr.cbChId, + handler: handler, + } + c.channelById[channelId] = info + + // Provide send callback if handler supports it + if setter, ok := handler.(interface{ SetSendFunc(func([]byte)) }); ok { + chId := channelId + setter.SetSendFunc(func(data []byte) { + c.SendDvcData(chId, data) + }) + } + slog.Debug("dvc: handler registered", "channel", channelName, "id", channelId) + } + + // If explicitly rejected, send a non-zero CreationStatus so the server + // does not use this channel (e.g. AUDIO_PLAYBACK_LOSSY_DVC → fallback to PCM). + if c.rejectedChannels[channelName] { + slog.Debug("dvc: rejecting channel", "channel", channelName, "id", channelId) + rspHdr := &DvcHeader{cmd: DYNVC_CREATE_REQ, sp: 0, cbChId: hdr.cbChId} + b := &bytes.Buffer{} + b.Write(rspHdr.serialize(channelId)) + core.WriteUInt32LE(0x80004005, b) // E_FAIL + c.Send(b.Bytes()) + return + } + + // Send success response (Sp SHOULD be 0 per MS-RDPEDYC 2.2.2.2). + // Always accept: some Windows servers stop sending data on static + // virtual channels (e.g. cliprdr) when DVC creation requests are + // rejected, even for unrelated channels. + rspHdr := &DvcHeader{cmd: DYNVC_CREATE_REQ, sp: 0, cbChId: hdr.cbChId} + b := &bytes.Buffer{} + b.Write(rspHdr.serialize(channelId)) + core.WriteUInt32LE(0, b) + c.Send(b.Bytes()) + + // Notify handler that channel is ready (CREATE_RSP has been sent) + if handler != nil { + if ch, ok := handler.(interface{ OnChannelCreated() }); ok { + ch.OnChannelCreated() + } + } +} + +func readDvcId(r io.Reader, cbLen uint8) (id uint32) { + switch cbLen { + case 0: + i, _ := core.ReadUInt8(r) + id = uint32(i) + case 1: + i, _ := core.ReadUint16LE(r) + id = uint32(i) + default: + id, _ = core.ReadUInt32LE(r) + } + return +} +func (c *DvcClient) processDataFirst(hdr *DvcHeader, s []byte) { + defer func() { + if r := recover(); r != nil { + slog.Error("dvc: panic in processDataFirst", "err", r) + } + }() + r := bytes.NewReader(s) + channelId := readDvcId(r, hdr.cbChId) + + // Read total length (encoding based on sp/Len field) + var totalLen uint32 + switch hdr.sp { + case 0: + l, _ := core.ReadUInt8(r) + totalLen = uint32(l) + case 1: + l, _ := core.ReadUint16LE(r) + totalLen = uint32(l) + default: + totalLen, _ = core.ReadUInt32LE(r) + } + + data, _ := core.ReadBytes(r.Len(), r) + ch, ok := c.channelById[channelId] + if !ok || ch.handler == nil { + return + } + + if uint32(len(data)) >= totalLen { + ch.handler.Process(data[:totalLen]) + } else { + ra := &dvcReassembly{totalLen: totalLen} + ra.buf.Write(data) + c.reassembly[channelId] = ra + } +} + +func (c *DvcClient) processData(hdr *DvcHeader, s []byte) { + defer func() { + if r := recover(); r != nil { + slog.Error("dvc: panic in processData", "err", r) + } + }() + r := bytes.NewReader(s) + channelId := readDvcId(r, hdr.cbChId) + data, _ := core.ReadBytes(r.Len(), r) + + ch, ok := c.channelById[channelId] + if !ok || ch.handler == nil { + return + } + + ra, hasReassembly := c.reassembly[channelId] + if hasReassembly { + ra.buf.Write(data) + if uint32(ra.buf.Len()) >= ra.totalLen { + ch.handler.Process(ra.buf.Bytes()[:ra.totalLen]) + delete(c.reassembly, channelId) + } + } else { + ch.handler.Process(data) + } +} + +func (c *DvcClient) processCapsPdu(hdr *DvcHeader, s []byte) { + r := bytes.NewReader(s) + core.ReadUInt8(r) + ver, _ := core.ReadUint16LE(r) + slog.Debug("Server supports dvc", "version", ver) + + // Respond with the server's version (up to 3). + // Version 3 is required for some servers to activate RDPGFX. + ver = min(ver, 3) + + // Client CAPS response: header(1) + pad(1) + version(2) = 4 bytes + // Priority charges are only in the server's CAPS request, not the client response. + b := &bytes.Buffer{} + core.WriteUInt8(0x50, b) // header: Cmd=5(CAPS), Sp=0, CbChId=0 + core.WriteUInt8(0x00, b) // pad + core.WriteUInt16LE(ver, b) + slog.Debug("dvc: CAPS response", "version", ver, "len", b.Len()) + c.Send(b.Bytes()) + c.negotiatedVersion = ver +} + +func (c *DvcClient) processSoftSyncRequest(hdr *DvcHeader, s []byte) { + r := bytes.NewReader(s) + core.ReadUInt8(r) // Pad + length, _ := core.ReadUInt32LE(r) // Length + flags, _ := core.ReadUint16LE(r) // Flags + numTunnels, _ := core.ReadUint16LE(r) + slog.Debug("DYNVC_SOFT_SYNC_REQUEST", "length", length, "flags", flags, "numTunnels", numTunnels) + + // Send SOFT_SYNC_RESPONSE: header + pad + length(4) + b := &bytes.Buffer{} + core.WriteUInt8((DYNVC_SOFT_SYNC_RESPONSE<<4)|0x00, b) // cmd=9, sp=0, cbChId=0 + core.WriteUInt8(0, b) // Pad + core.WriteUInt32LE(0x04, b) // Length = 4 (just the length field) + c.Send(b.Bytes()) +} diff --git a/plugin/rail/rail.go b/plugin/rail/rail.go new file mode 100644 index 0000000..d6f499f --- /dev/null +++ b/plugin/rail/rail.go @@ -0,0 +1,452 @@ +// rail.go +package rail + +import ( + "bytes" + "encoding/hex" + "fmt" + "log/slog" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/plugin" +) + +const ( + ChannelName = plugin.RAIL_SVC_CHANNEL_NAME + ChannelOption = plugin.CHANNEL_OPTION_INITIALIZED | plugin.CHANNEL_OPTION_ENCRYPT_RDP | + plugin.CHANNEL_OPTION_COMPRESS_RDP | plugin.CHANNEL_OPTION_SHOW_PROTOCOL +) + +const ( + TS_RAIL_ORDER_EXEC = 0x0001 + TS_RAIL_ORDER_ACTIVATE = 0x0002 + TS_RAIL_ORDER_SYSPARAM = 0x0003 + TS_RAIL_ORDER_SYSCOMMAND = 0x0004 + TS_RAIL_ORDER_HANDSHAKE = 0x0005 + TS_RAIL_ORDER_NOTIFY_EVENT = 0x0006 + TS_RAIL_ORDER_WINDOWMOVE = 0x0008 + TS_RAIL_ORDER_LOCALMOVESIZE = 0x0009 + TS_RAIL_ORDER_MINMAXINFO = 0x000A + TS_RAIL_ORDER_CLIENTSTATUS = 0x000B + TS_RAIL_ORDER_SYSMENU = 0x000C + TS_RAIL_ORDER_LANGBARINFO = 0x000D + TS_RAIL_ORDER_GET_APPID_REQ = 0x000E + TS_RAIL_ORDER_GET_APPID_RESP = 0x000F + TS_RAIL_ORDER_TASKBARINFO = 0x0010 + TS_RAIL_ORDER_LANGUAGEIMEINFO = 0x0011 + TS_RAIL_ORDER_COMPARTMENTINFO = 0x0012 + TS_RAIL_ORDER_HANDSHAKE_EX = 0x0013 + TS_RAIL_ORDER_ZORDER_SYNC = 0x0014 + TS_RAIL_ORDER_CLOAK = 0x0015 + TS_RAIL_ORDER_POWER_DISPLAY_REQUEST = 0x0016 + TS_RAIL_ORDER_SNAP_ARRANGE = 0x0017 + TS_RAIL_ORDER_GET_APPID_RESP_EX = 0x0018 + TS_RAIL_ORDER_EXEC_RESULT = 0x0080 +) + +type RailClient struct { + w core.ChannelSender + DesktopWidth uint16 + DesktopHeight uint16 + RemoteApplicationProgram string + ShellWorkingDirectory string + RemoteApplicationCmdLine string +} + +func NewClient() *RailClient { + return &RailClient{ + DesktopWidth: 800, + DesktopHeight: 600, + RemoteApplicationProgram: "calc", + ShellWorkingDirectory: "/tmp", + } +} + +type RailPDUHeader struct { + OrderType uint16 `struc:"little"` + OrderLength uint16 `struc:"little"` +} + +func NewRailPDUHeader(mType, ln uint16) *RailPDUHeader { + return &RailPDUHeader{ + OrderType: mType, + OrderLength: ln, + } +} + +func (h *RailPDUHeader) serialize() []byte { + b := &bytes.Buffer{} + core.WriteUInt16LE(h.OrderType, b) + core.WriteUInt16LE(h.OrderLength, b) + return b.Bytes() +} + +func (c *RailClient) sendData(mType uint16, ln int, s []byte) { + slog.Debug("sendData", "ln", ln, "s_length", len(s), "data", hex.EncodeToString(s)) + header := NewRailPDUHeader(mType, uint16(ln)) + + b := &bytes.Buffer{} + core.WriteBytes(header.serialize(), b) + core.WriteBytes(s, b) + + c.Send(b.Bytes()) +} + +func (c *RailClient) Send(s []byte) (int, error) { + slog.Debug("send", "len", len(s), "data", hex.EncodeToString(s)) + name, _ := c.GetType() + return c.w.SendToChannel(name, s) +} +func (c *RailClient) Sender(f core.ChannelSender) { + c.w = f +} +func (c *RailClient) GetType() (string, uint32) { + return ChannelName, ChannelOption +} + +func (c *RailClient) Process(s []byte) { + slog.Debug("recv", "data", hex.EncodeToString(s)) + r := bytes.NewReader(s) + msgType, _ := core.ReadUint16LE(r) + length, _ := core.ReadUint16LE(r) + + slog.Debug("rail", "type", fmt.Sprintf("0x%x", msgType), "length", length, "remaining", r.Len()) + + b, _ := core.ReadBytes(int(length), r) + slog.Debug("recv body", "data", hex.EncodeToString(b)) + + switch msgType { + case TS_RAIL_ORDER_HANDSHAKE: + slog.Debug("TS_RAIL_ORDER_HANDSHAKE") + c.processOrderHandshake(b) + case TS_RAIL_ORDER_SYSPARAM: + slog.Debug("TS_RAIL_ORDER_SYSPARAM") + c.processOrderSysparam(b) + case TS_RAIL_ORDER_EXEC_RESULT: + slog.Debug("TS_RAIL_ORDER_EXEC_RESULT") + c.processExecResult(b) + + default: + slog.Error("type not supported", "msgType", fmt.Sprintf("0x%x", msgType)) + } +} + +func (c *RailClient) processOrderHandshake(b []byte) { + r := bytes.NewReader(b) + buildNumber, _ := core.ReadUInt32LE(r) + slog.Debug("processOrderHandshake", "buildNumber", buildNumber) + + //send client info + c.sendClientStatus() + + //send client systemparam + c.sendClientSystemparam() + + //send client execute + c.sendClientExecute() +} + +const ( + TS_RAIL_CLIENTSTATUS_ALLOWLOCALMOVESIZE = 0x00000001 + TS_RAIL_CLIENTSTATUS_AUTORECONNECT = 0x00000002 + TS_RAIL_CLIENTSTATUS_ZORDER_SYNC = 0x00000004 + TS_RAIL_CLIENTSTATUS_WINDOW_RESIZE_MARGIN_SUPPORTED = 0x00000010 + TS_RAIL_CLIENTSTATUS_HIGH_DPI_ICONS_SUPPORTED = 0x00000020 + TS_RAIL_CLIENTSTATUS_APPBAR_REMOTING_SUPPORTED = 0x00000040 + TS_RAIL_CLIENTSTATUS_POWER_DISPLAY_REQUEST_SUPPORTED = 0x00000080 + TS_RAIL_CLIENTSTATUS_GET_APPID_RESPONSE_EX_SUPPORTED = 0x00000100 + TS_RAIL_CLIENTSTATUS_BIDIRECTIONAL_CLOAK_SUPPORTED = 0x00000200 +) + +func (c *RailClient) sendClientStatus() { + slog.Debug("Send client Status") + var flags uint32 = TS_RAIL_CLIENTSTATUS_ALLOWLOCALMOVESIZE + + //if (settings->AutoReconnectionEnabled) + //clientStatus.flags |= TS_RAIL_CLIENTSTATUS_AUTORECONNECT; + + flags |= TS_RAIL_CLIENTSTATUS_ZORDER_SYNC + flags |= TS_RAIL_CLIENTSTATUS_WINDOW_RESIZE_MARGIN_SUPPORTED + flags |= TS_RAIL_CLIENTSTATUS_APPBAR_REMOTING_SUPPORTED + flags |= TS_RAIL_CLIENTSTATUS_POWER_DISPLAY_REQUEST_SUPPORTED + flags |= TS_RAIL_CLIENTSTATUS_BIDIRECTIONAL_CLOAK_SUPPORTED + + b := &bytes.Buffer{} + core.WriteUInt32LE(flags, b) + + c.sendData(TS_RAIL_ORDER_CLIENTSTATUS, 4, b.Bytes()) +} + +const ( + SPI_SET_SCREEN_SAVE_ACTIVE = 0x00000011 + SPI_SET_SCREEN_SAVE_SECURE = 0x00000077 +) +const ( + /*Bit mask values for SPI_ parameters*/ + SPI_MASK_SET_DRAG_FULL_WINDOWS = 0x00000001 + SPI_MASK_SET_KEYBOARD_CUES = 0x00000002 + SPI_MASK_SET_KEYBOARD_PREF = 0x00000004 + SPI_MASK_SET_MOUSE_BUTTON_SWAP = 0x00000008 + SPI_MASK_SET_WORK_AREA = 0x00000010 + SPI_MASK_DISPLAY_CHANGE = 0x00000020 + SPI_MASK_TASKBAR_POS = 0x00000040 + SPI_MASK_SET_HIGH_CONTRAST = 0x00000080 + SPI_MASK_SET_SCREEN_SAVE_ACTIVE = 0x00000100 + SPI_MASK_SET_SET_SCREEN_SAVE_SECURE = 0x00000200 + SPI_MASK_SET_CARET_WIDTH = 0x00000400 + SPI_MASK_SET_STICKY_KEYS = 0x00000800 + SPI_MASK_SET_TOGGLE_KEYS = 0x00001000 + SPI_MASK_SET_FILTER_KEYS = 0x00002000 +) +const ( + SPI_SET_DRAG_FULL_WINDOWS = 0x00000025 + SPI_SET_KEYBOARD_CUES = 0x0000100B + SPI_SET_KEYBOARD_PREF = 0x00000045 + SPI_SET_MOUSE_BUTTON_SWAP = 0x00000021 + SPI_SET_WORK_AREA = 0x0000002F + SPI_DISPLAY_CHANGE = 0x0000F001 + SPI_TASKBAR_POS = 0x0000F000 + SPI_SET_HIGH_CONTRAST = 0x00000043 + SPI_SETCARETWIDTH = 0x00002007 + SPI_SETSTICKYKEYS = 0x0000003B + SPI_SETTOGGLEKEYS = 0x00000035 + SPI_SETFILTERKEYS = 0x00000033 +) + +type TsFilterKeys struct { + Flags uint32 + WaitTime uint32 + DelayTime uint32 + RepeatTime uint32 + BounceTime uint32 +} +type RailHighContrast struct { + flags uint32 + colorSchemeLength uint32 + colorScheme string +} +type Rectangle16 struct { + left uint16 + top uint16 + right uint16 + bottom uint16 +} +type RailSysparamOrder struct { + param uint32 + params uint32 + dragFullWindows uint8 + keyboardCues uint8 + keyboardPref uint8 + mouseButtonSwap uint8 + workArea Rectangle16 + displayChange Rectangle16 + taskbarPos Rectangle16 + highContrast RailHighContrast + caretWidth uint32 + stickyKeys uint32 + toggleKeys uint32 + filterKeys TsFilterKeys + setScreenSaveActive uint8 + setScreenSaveSecure uint8 +} + +func (c *RailClient) sendClientSystemparam() { + slog.Debug("Send client Systemparam") + + var sp RailSysparamOrder + sp.params = 0 + sp.params |= SPI_MASK_SET_HIGH_CONTRAST + sp.highContrast.colorScheme = "" + sp.highContrast.colorSchemeLength = 0 + sp.highContrast.flags = 0x7E + sp.params |= SPI_MASK_SET_MOUSE_BUTTON_SWAP + sp.mouseButtonSwap = 0 + sp.params |= SPI_MASK_SET_KEYBOARD_PREF + sp.keyboardPref = 0 + sp.params |= SPI_MASK_SET_DRAG_FULL_WINDOWS + sp.dragFullWindows = 0 + sp.params |= SPI_MASK_SET_KEYBOARD_CUES + sp.keyboardCues = 0 + sp.params |= SPI_MASK_SET_WORK_AREA + sp.workArea.left = 0 + sp.workArea.top = 0 + sp.workArea.right = c.DesktopWidth + sp.workArea.bottom = c.DesktopHeight + + if sp.params&SPI_MASK_SET_HIGH_CONTRAST != 0 { + sp.param = SPI_SET_HIGH_CONTRAST + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_TASKBAR_POS != 0 { + sp.param = SPI_TASKBAR_POS + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_SET_MOUSE_BUTTON_SWAP != 0 { + sp.param = SPI_SET_MOUSE_BUTTON_SWAP + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_SET_KEYBOARD_PREF != 0 { + sp.param = SPI_SET_KEYBOARD_PREF + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_SET_DRAG_FULL_WINDOWS != 0 { + sp.param = SPI_SET_DRAG_FULL_WINDOWS + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_SET_KEYBOARD_CUES != 0 { + sp.param = SPI_SET_KEYBOARD_CUES + c.sendOneClientSysparam(&sp) + } + + if sp.params&SPI_MASK_SET_WORK_AREA != 0 { + sp.param = SPI_SET_WORK_AREA + slog.Debug("SPI_SET_WORK_AREA") + c.sendOneClientSysparam(&sp) + } +} + +func (c *RailClient) sendOneClientSysparam(sp *RailSysparamOrder) { + length := 0 + b := &bytes.Buffer{} + core.WriteUInt32LE(sp.param, b) + switch sp.param { + case SPI_SET_DRAG_FULL_WINDOWS: + core.WriteUInt8(sp.dragFullWindows, b) + + case SPI_SET_KEYBOARD_CUES: + core.WriteUInt8(sp.keyboardCues, b) + + case SPI_SET_KEYBOARD_PREF: + core.WriteUInt8(sp.keyboardPref, b) + + case SPI_SET_MOUSE_BUTTON_SWAP: + core.WriteUInt8(sp.mouseButtonSwap, b) + + case SPI_SET_WORK_AREA: + core.WriteUInt16LE(sp.workArea.left, b) + core.WriteUInt16LE(sp.workArea.top, b) + core.WriteUInt16LE(sp.workArea.right, b) + core.WriteUInt16LE(sp.workArea.bottom, b) + + case SPI_DISPLAY_CHANGE: + core.WriteUInt16LE(sp.displayChange.left, b) + core.WriteUInt16LE(sp.displayChange.top, b) + core.WriteUInt16LE(sp.displayChange.right, b) + core.WriteUInt16LE(sp.displayChange.bottom, b) + + case SPI_TASKBAR_POS: + core.WriteUInt16LE(sp.taskbarPos.left, b) + core.WriteUInt16LE(sp.taskbarPos.top, b) + core.WriteUInt16LE(sp.taskbarPos.right, b) + core.WriteUInt16LE(sp.taskbarPos.bottom, b) + + case SPI_SET_HIGH_CONTRAST: + core.WriteUInt32LE(sp.highContrast.flags, b) + core.WriteUInt32LE(sp.highContrast.colorSchemeLength, b) + data := core.UnicodeEncode(sp.highContrast.colorScheme) + core.WriteBytes(data, b) + + case SPI_SETFILTERKEYS: + core.WriteUInt32LE(sp.filterKeys.Flags, b) + core.WriteUInt32LE(sp.filterKeys.WaitTime, b) + core.WriteUInt32LE(sp.filterKeys.DelayTime, b) + core.WriteUInt32LE(sp.filterKeys.RepeatTime, b) + core.WriteUInt32LE(sp.filterKeys.BounceTime, b) + + case SPI_SETSTICKYKEYS: + core.WriteUInt32LE(sp.stickyKeys, b) + + case SPI_SETCARETWIDTH: + core.WriteUInt32LE(sp.caretWidth, b) + + case SPI_SETTOGGLEKEYS: + core.WriteUInt32LE(sp.toggleKeys, b) + + case SPI_MASK_SET_SET_SCREEN_SAVE_SECURE: + core.WriteUInt8(sp.setScreenSaveSecure, b) + + case SPI_MASK_SET_SCREEN_SAVE_ACTIVE: + core.WriteUInt8(sp.setScreenSaveActive, b) + + default: + slog.Error("ERROR_BAD_ARGUMENTS") + return + } + + c.sendData(TS_RAIL_ORDER_SYSPARAM, length+b.Len(), b.Bytes()) +} + +type RailExecOrder struct { + flags uint16 + RemoteApplicationProgram string + RemoteApplicationWorkingDir string + RemoteApplicationArguments string +} + +func (c *RailClient) sendClientExecute() { + slog.Debug("Send Client Execute") + var exec RailExecOrder + //exec.flags = TS_RAIL_EXEC_FLAG_EXPAND_ARGUMENTS + exec.RemoteApplicationProgram = c.RemoteApplicationProgram + exec.RemoteApplicationWorkingDir = c.ShellWorkingDirectory + exec.RemoteApplicationArguments = c.RemoteApplicationCmdLine + + program := core.UnicodeEncode(exec.RemoteApplicationProgram) + workdir := core.UnicodeEncode(exec.RemoteApplicationWorkingDir) + arguments := core.UnicodeEncode(exec.RemoteApplicationArguments) + + length := 4 + b := &bytes.Buffer{} + core.WriteUInt16LE(exec.flags, b) + core.WriteUInt16LE(uint16(len(program)), b) + core.WriteUInt16LE(uint16(len(workdir)), b) + core.WriteUInt16LE(uint16(len(arguments)), b) + core.WriteBytes(program, b) + core.WriteBytes(workdir, b) + core.WriteBytes(arguments, b) + length += b.Len() + + c.sendData(TS_RAIL_ORDER_EXEC, length, b.Bytes()) + +} + +func (c *RailClient) processOrderSysparam(b []byte) { + r := bytes.NewReader(b) + systemParam, _ := core.ReadUInt32LE(r) + body, _ := core.ReadUInt8(r) + slog.Debug("processOrderSysparam", "systemParam", fmt.Sprintf("0x%x", systemParam), "body", body) +} + +const ( + //The Client Execute request was successful and the requested application or file has been launched. + RAIL_EXEC_S_OK = 0x0000 + //The Client Execute request could not be satisfied because the server is not monitoring the current input desktop. + RAIL_EXEC_E_HOOK_NOT_LOADED = 0x0001 + //The Execute request could not be satisfied because the request PDU was malformed. + RAIL_EXEC_E_DECODE_FAILED = 0x0002 + //The Client Execute request could not be satisfied because the requested application was blocked by policy from being launched on the server. + RAIL_EXEC_E_NOT_IN_ALLOWLIST = 0x0003 + //The Client Execute request could not be satisfied because the application or file path could not be found. + RAIL_EXEC_E_FILE_NOT_FOUND = 0x0005 + //The Client Execute request could not be satisfied because an unspecified error occurred on the server. + RAIL_EXEC_E_FAIL = 0x0006 + //The Client Execute request could not be satisfied because the remote session is locked. + RAIL_EXEC_E_SESSION_LOCKED = 0x0007 +) + +func (c *RailClient) processExecResult(b []byte) { + r := bytes.NewReader(b) + flags, _ := core.ReadUint16LE(r) + execResult, _ := core.ReadUint16LE(r) + rawResult, _ := core.ReadUInt32LE(r) + core.ReadUint16LE(r) + exeOrFileLength, _ := core.ReadUint16LE(r) + exeOrFile, _ := core.ReadBytes(r.Len(), r) + slog.Debug("processExecResult", "flags", flags, "execResult", execResult, "rawResult", rawResult) + slog.Debug("processExecResult", "length", exeOrFileLength, "file", core.UnicodeDecode(exeOrFile)) +} diff --git a/plugin/rdpdr/drive.go b/plugin/rdpdr/drive.go new file mode 100644 index 0000000..a88f509 --- /dev/null +++ b/plugin/rdpdr/drive.go @@ -0,0 +1,256 @@ +// drive.go — MS-RDPEFS 设备 I/O 响应的编码器:目录枚举条目 +// ([MS-FSCC] FILE_*_DIRECTORY_INFORMATION)、文件/卷信息结构与 UTF-16 +// 辅助。时间统一从 JS 侧的毫秒时间戳换算 Windows FILETIME。 +package rdpdr + +import ( + "encoding/binary" + "encoding/json" + "unicode/utf16" +) + +// DirEntry 是桥接侧返回的单条目录/文件元数据。 +type DirEntry struct { + Name string `json:"name"` + Dir bool `json:"dir"` + Size int64 `json:"size"` + Mtime int64 `json:"mtime"` // 毫秒时间戳(最后写入) + Created int64 `json:"created"` // 毫秒时间戳(缺省取 Mtime) + Accessed int64 `json:"accessed"` // 毫秒时间戳(缺省取 Mtime) +} + +// FSCC 信息类(本实现支持的子集)。 +const ( + FileDirectoryInformation = 0x00000001 + FileFullDirectoryInformation = 0x00000002 + FileBothDirectoryInformation = 0x00000003 + FileNamesInformation = 0x0000000C + + FileFsVolumeInformation = 0x00000001 + FileFsSizeInformation = 0x00000003 + FileFsDeviceInformation = 0x00000004 + FileFsAttributeInformation = 0x00000005 + FileFsFullSizeInformation = 0x00000007 + + FileBasicInformation = 0x00000004 + FileStandardInformation = 0x00000005 + FileNetworkOpenInformation = 0x00000022 +) + +// filetime 把毫秒 Unix 时间戳换算为 Windows FILETIME(100ns,1601 纪元)。 +func filetime(ms int64) uint64 { + const epochDelta = 116444736000000000 // 1601→1970 的 100ns 数 + if ms < 0 { + ms = 0 + } + return uint64(ms)*10000 + epochDelta +} + +func utf16Bytes(s string) []byte { + u := utf16.Encode([]rune(s)) + b := make([]byte, 2*len(u)) + for i, v := range u { + binary.LittleEndian.PutUint16(b[2*i:], v) + } + return b +} + +func utf16ToString(b []byte) string { + if len(b)%2 != 0 { + b = b[:len(b)-1] + } + u := make([]uint16, len(b)/2) + for i := range u { + u[i] = binary.LittleEndian.Uint16(b[2*i:]) + } + runes := utf16.Decode(u) + // 去掉结尾 NUL + for n := len(runes) - 1; n >= 0; n-- { + if runes[n] == 0 { + runes = runes[:n] + } else { + break + } + } + return string(runes) +} + +// decodeEntries 解析桥接 JSON 数组。 +func decodeEntries(jsonStr string) ([]DirEntry, error) { + var entries []DirEntry + if err := json.Unmarshal([]byte(jsonStr), &entries); err != nil { + return nil, err + } + for i := range entries { + fillTimes(&entries[i]) + } + return entries, nil +} + +// decodeEntry 解析桥接 JSON 单对象。 +func decodeEntry(jsonStr string) (DirEntry, error) { + var e DirEntry + if err := json.Unmarshal([]byte(jsonStr), &e); err != nil { + return e, err + } + fillTimes(&e) + return e, nil +} + +func fillTimes(e *DirEntry) { + if e.Created == 0 { + e.Created = e.Mtime + } + if e.Accessed == 0 { + e.Accessed = e.Mtime + } +} + +func (e *DirEntry) attributes() uint32 { + if e.Dir { + return FILE_ATTRIBUTE_DIRECTORY + } + if e.Size == 0 { + return FILE_ATTRIBUTE_NORMAL + } + return FILE_ATTRIBUTE_NORMAL | FILE_ATTRIBUTE_ARCHIVE +} + +// encodeDirEntries 把一批条目编码为 FSCC 目录信息结构链。 +// 支持 FileDirectoryInformation(1)/FileBothDirectoryInformation(3)/ +// FileNamesInformation(0xC);其它返回 nil(上层回 NOT_IMPLEMENTED)。 +func encodeDirEntries(infoClass uint32, entries []DirEntry) []byte { + if len(entries) == 0 { + return []byte{} + } + var out []byte + for i := range entries { + e := &entries[i] + name := utf16Bytes(e.Name) + var raw []byte + switch infoClass { + case FileDirectoryInformation: + // NextEntryOffset(4) FileIndex(4) Creation(8) Access(8) Write(8) + // Change(8) EndOfFile(8) Allocation(8) Attributes(4) NameLen(4) = 64 + raw = make([]byte, 64+len(name)) + putTimes(raw, 8, e) + binary.LittleEndian.PutUint64(raw[40:], uint64(e.Size)) + binary.LittleEndian.PutUint64(raw[48:], uint64(e.Size)) + binary.LittleEndian.PutUint32(raw[56:], e.attributes()) + binary.LittleEndian.PutUint32(raw[60:], uint32(len(name))) + copy(raw[64:], name) + case FileBothDirectoryInformation: + // 64 字节同上 + EaSize(4) ShortNameLen(1) ShortName(24) = 93 + raw = make([]byte, 93+len(name)) + putTimes(raw, 8, e) + binary.LittleEndian.PutUint64(raw[40:], uint64(e.Size)) + binary.LittleEndian.PutUint64(raw[48:], uint64(e.Size)) + binary.LittleEndian.PutUint32(raw[56:], e.attributes()) + binary.LittleEndian.PutUint32(raw[60:], uint32(len(name))) + raw[68] = 0 // ShortNameLength + // ShortName[24] 全 0(无短名) + copy(raw[93:], name) + case FileNamesInformation: + // NextEntryOffset(4) FileIndex(4) FileNameLength(4) = 12 + raw = make([]byte, 12+len(name)) + binary.LittleEndian.PutUint32(raw[8:], uint32(len(name))) + copy(raw[12:], name) + default: + return nil + } + if i == len(entries)-1 { + // 末条:NextEntryOffset=0,无需对齐 + binary.LittleEndian.PutUint32(raw[0:], 0) + out = append(out, raw...) + continue + } + // 非末条:长度 4 字节对齐,补零 + next := (len(raw) + 3) &^ 3 + buf := make([]byte, next) + copy(buf, raw) + binary.LittleEndian.PutUint32(buf[0:], uint32(next)) + out = append(out, buf...) + } + return out +} + +// putTimes 在 off 处写 Creation/Access/Write/Change 四个 FILETIME(32 字节)。 +func putTimes(b []byte, off int, e *DirEntry) { + binary.LittleEndian.PutUint64(b[off:], filetime(e.Created)) + binary.LittleEndian.PutUint64(b[off+8:], filetime(e.Accessed)) + binary.LittleEndian.PutUint64(b[off+16:], filetime(e.Mtime)) + binary.LittleEndian.PutUint64(b[off+24:], filetime(e.Mtime)) +} + +// encodeFileInfo 编码 QUERY_INFORMATION 响应体(不含 Length 前缀)。 +func encodeFileInfo(infoClass uint32, e DirEntry) []byte { + switch infoClass { + case FileBasicInformation: + // Creation(8) Access(8) Write(8) Change(8) Attributes(4) Reserved(4) = 40 + b := make([]byte, 40) + putTimes(b, 0, &e) + binary.LittleEndian.PutUint32(b[32:], e.attributes()) + return b + case FileStandardInformation: + // AllocationSize(8) EndOfFile(8) NumberOfLinks(4) DeletePending(1) Directory(1) = 22 + b := make([]byte, 22) + binary.LittleEndian.PutUint64(b[0:], uint64(e.Size)) + binary.LittleEndian.PutUint64(b[8:], uint64(e.Size)) + binary.LittleEndian.PutUint32(b[16:], 1) + b[20] = 0 + if e.Dir { + b[21] = 1 + } + return b + case FileNetworkOpenInformation: + // Creation(8) Access(8) Write(8) Change(8) Allocation(8) EndOfFile(8) Attributes(4) = 56 + b := make([]byte, 56) + putTimes(b, 0, &e) + binary.LittleEndian.PutUint64(b[32:], uint64(e.Size)) + binary.LittleEndian.PutUint64(b[40:], uint64(e.Size)) + binary.LittleEndian.PutUint32(b[48:], e.attributes()) + return b + default: + return nil + } +} + +// encodeVolumeInfo 编码 QUERY_VOLUME_INFORMATION 响应体(不含 Length 前缀)。 +func encodeVolumeInfo(infoClass uint32, label string) []byte { + switch infoClass { + case FileFsVolumeInformation: + // VolumeCreationTime(8) SerialNumber(4) LabelLength(4) SupportsObjects(1) Label + vol := utf16Bytes(label) + b := make([]byte, 17+len(vol)) + binary.LittleEndian.PutUint32(b[8:], 0x1ABCF2D8) // 任意固定序列号 + binary.LittleEndian.PutUint32(b[12:], uint32(len(vol))) + copy(b[17:], vol) + return b + case FileFsSizeInformation, FileFsFullSizeInformation: + // TotalAllocationUnits(8) Available(8) SectorsPerUnit(4) BytesPerSector(4) = 24 + // (Win10 服务器挂载设备时常探测 FullSize——缺失会致设备"不支持") + b := make([]byte, 24) + binary.LittleEndian.PutUint64(b[0:], 0x00100000) + binary.LittleEndian.PutUint64(b[8:], 0x00080000) + binary.LittleEndian.PutUint32(b[16:], 8) + binary.LittleEndian.PutUint32(b[20:], 512) + return b + case FileFsDeviceInformation: + // DeviceType(4)=FILE_DEVICE_DISK Characteristics(4) = 8 + b := make([]byte, 8) + binary.LittleEndian.PutUint32(b[0:], 7) + return b + case FileFsAttributeInformation: + // Attributes(4) MaxComponentLen(4) NameLength(4) Name("FAT32",规避 + // 服务端按 NTFS 语义发起的 ACL/重解析点等操作——mstsc/FreeRDP 同款选择) + fs := utf16Bytes("FAT32") + b := make([]byte, 12+len(fs)) + binary.LittleEndian.PutUint32(b[0:], 0x00000007) // CASE_SENSITIVE_SEARCH|CASE_PRESERVED_NAMES|UNICODE_ON_DISK + binary.LittleEndian.PutUint32(b[4:], 255) + binary.LittleEndian.PutUint32(b[8:], uint32(len(fs))) + copy(b[12:], fs) + return b + default: + return nil + } +} diff --git a/plugin/rdpdr/rdpdr.go b/plugin/rdpdr/rdpdr.go new file mode 100644 index 0000000..7f7590c --- /dev/null +++ b/plugin/rdpdr/rdpdr.go @@ -0,0 +1,695 @@ +// Package rdpdr implements the client side of the RDP File System Virtual +// Channel Extension (MS-RDPEFS, channel name "rdpdr") for drive redirection: +// the server sees a file-system device backed by an asynchronous bridge +// (browser-picked folder), and enumerates/reads it through device I/O +// requests. +// +// The bridge is asynchronous by necessity (browser filesystem operations are +// promise-based): each Device I/O Request is dispatched to the Filesystem +// implementation together with its completionId, and the result — or failure — +// is fed back via CompleteStatus/CompleteBytes/CompleteJSON, which emit the +// matching Device I/O Response with the proper wire encoding. +package rdpdr + +import ( + "encoding/binary" + "log/slog" + "strings" + "sync" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +// ── 常量(MS-RDPEFS 2.2,与 FreeRDP channels/rdpdr.h 对齐)──────────────── + +const ( + RDPDR_CTYP_CORE = 0x4472 + RDPDR_CTYP_PRN = 0x5052 +) + +const ( + PAKID_CORE_SERVER_ANNOUNCE = 0x496E + PAKID_CORE_CLIENTID_CONFIRM = 0x4343 // 双向:客户端 Announce Reply / 服务端 Confirm + PAKID_CORE_CLIENT_NAME = 0x434E + PAKID_CORE_DEVICELIST_ANNOUNCE = 0x4441 + PAKID_CORE_DEVICE_REPLY = 0x6472 + PAKID_CORE_DEVICE_IOREQUEST = 0x4952 + PAKID_CORE_DEVICE_IOCOMPLETION = 0x4943 + PAKID_CORE_SERVER_CAPABILITY = 0x5350 + PAKID_CORE_CLIENT_CAPABILITY = 0x4350 + PAKID_CORE_DEVICELIST_REMOVE = 0x444D + PAKID_CORE_USER_LOGGEDON = 0x554C +) + +const ( + CAP_GENERAL_TYPE = 0x0001 + CAP_DRIVE_TYPE = 0x0004 +) + +// 协议版本与能力位(与 FreeRDP channels/rdpdr.h / rdpdr_capabilities.c 对齐)。 +const ( + RDPDR_VERSION_MINOR_RDP51 = 0x0005 + RDPDR_VERSION_MINOR_RDP10X = 0x000D // 客户端响应版本上限(FreeRDP rdpdr_main.c MIN 上限) + + // GENERAL capset 的 ExtendedPDU 能力位(MS-RDPEFS 2.2.2.1) + RDPDR_DEVICE_REMOVE_PDUS = 0x00000001 + RDPDR_CLIENT_DISPLAY_NAME_PDU = 0x00000002 + RDPDR_USER_LOGGEDON_PDU = 0x00000004 + // GENERAL capset 的 extraFlags1 + RDPDR_ENABLE_ASYNCIO = 0x00000001 +) + +// clientIOCode1 是能力响应里宣告的 IRP major 码位掩码——取 FreeRDP 同款 +// 全集(含本实现未细分的 CLEANUP/FLUSH/SHUTDOWN/LOCK/SECURITY 等,这些 +// 会以 STATUS_NOT_IMPLEMENTED 兜底完成),与服务端 ioCode1 求交后回给服务端。 +const clientIOCode1 = 1< 的 RDPNP 共享映射——访问报"无法访问"且零通道 IO, + // 见 doc/history/stage6-plan.md(RDPDR-1/2 章节)。 + clientIDConfirmed bool + deviceListSent bool + + fs Filesystem + + nextFileID uint32 + pendingMu sync.Mutex + pending map[uint32]*pending // completionId → pending + + enumMu sync.Mutex + enumPos map[uint32]int // FileId → 已返回条目数(目录枚举分页游标) +} + +// clientComputerName 是 CLIENT_NAME_REQUEST 里上报的客户端机器名。取值 +// "tsclient" 与 M1 端到端验收通过(2026-09-12 08:25,10.0.0.3)时的取值 +// 一致;后改为 "grdpclient" 的会话全部失败。因当时 CLOSE 探测错误会独立 +// 导致映射被撤销,两个变量未分离验证,先回退到已知良好值(注释勿 +// overclaim:名字的独立影响待 CLOSE 修复验证后再做 A/B)。 +const clientComputerName = "tsclient" + +// NewHandler 创建处理器。shareName 是共享名:DEVICE_ANNOUNCE 的 DosName 取 +// 其前 8 字符(FreeRDP 同款),DeviceData 为 ASCII 全名 + NUL——服务端按 +// DosName 建 \\tsclient\ 的 UNC 映射。label 是卷标(可含中文)。 +func NewHandler(shareName, label string) *Handler { + shareName = sanitizeShareName(shareName) + return &Handler{ + devName: shareName, + label: label, + versionMinor: RDPDR_VERSION_MINOR_RDP10X, + deviceID: 1, // DeviceId 从 1 起(FreeRDP 同款;0 可能与服务端内部路由冲突) + nextFileID: 1, + pending: make(map[uint32]*pending), + enumPos: make(map[uint32]int), + } +} + +// SetFilesystem 挂接异步文件系统桥(连接前调用)。 +func (h *Handler) SetFilesystem(fs Filesystem) { h.fs = fs } + +// min16 返回较小者(Go 1.21 前无泛型 min,wasm 目标锁定旧工具链时需要)。 +func min16(a, b uint16) uint16 { + if a < b { + return a + } + return b +} + +// sanitizeShareName 把共享名中 DEVICE_ANNOUNCE 不允许的字符替换为 '_' +//(MS-RDPEFS 2.2.1.3;FreeRDP drive_main.c 同款过滤表,含空格与逗号)。 +func sanitizeShareName(s string) string { + const forbidden = `\/:*?"<>|, ` + "\t" + r := []rune(s) + for i, c := range r { + if strings.ContainsRune(forbidden, c) { + r[i] = '_' + } + } + return string(r) +} + +// SetDeviceID 指定宣告用的 DeviceId(默认 1)。 +func (h *Handler) SetDeviceID(id uint32) { h.deviceID = id } + +// GetType 实现 plugin.ChannelTransport。 +func (h *Handler) GetType() (string, uint32) { + return "rdpdr", 0x80000000 | 0x40000000 | 0x00800000 // INITIALIZED|ENCRYPT_RDP|COMPRESS_RDP +} + +// Sender 实现 plugin.ChannelTransport。 +func (h *Handler) Sender(s core.ChannelSender) { h.channelSender = s } + +func (h *Handler) send(b []byte) { + if h.channelSender == nil { + return + } + if _, err := h.channelSender.SendToChannel("rdpdr", b); err != nil { + slog.Warn("rdpdr send", "err", err) + } +} + +// Process 实现 plugin.ChannelTransport:分发服务端消息。 +func (h *Handler) Process(data []byte) { + if len(data) < 4 { + return + } + component := binary.LittleEndian.Uint16(data[0:]) + packetID := binary.LittleEndian.Uint16(data[2:]) + if component != RDPDR_CTYP_CORE { + slog.Debug("rdpdr: non-core component", "component", component, "packetID", packetID) + return + } + switch packetID { + case PAKID_CORE_SERVER_ANNOUNCE: + h.processServerAnnounce(data) + case PAKID_CORE_CLIENTID_CONFIRM: + // 服务端回显其接受的版本与 ClientId(FreeRDP 采纳该版本), + // 1.0x 语义下设备列表等 USER_LOGGEDON 再发 + if len(data) >= 12 { + h.versionMajor = binary.LittleEndian.Uint16(data[4:]) + h.versionMinor = binary.LittleEndian.Uint16(data[6:]) + h.clientID = binary.LittleEndian.Uint32(data[8:]) + } + slog.Debug("rdpdr: server clientid confirm", + "major", h.versionMajor, "minor", h.versionMinor, "clientID", h.clientID) + h.clientIDConfirmed = true + case PAKID_CORE_USER_LOGGEDON: + slog.Debug("rdpdr: user loggedon") + if h.clientIDConfirmed && !h.deviceListSent { + h.sendDeviceList() + } + case PAKID_CORE_DEVICE_REPLY: + if len(data) >= 12 { + slog.Debug("rdpdr: device reply", "deviceID", binary.LittleEndian.Uint32(data[4:]), + "result", binary.LittleEndian.Uint32(data[8:])) + } + case PAKID_CORE_SERVER_CAPABILITY: + h.processServerCapability(data) + case PAKID_CORE_DEVICE_IOREQUEST: + h.processIORequest(data) + default: + slog.Debug("rdpdr: unhandled", "packetID", packetID, "len", len(data)) + } +} + +// ── 握手:announce reply → name request → 设备列表 ─────────────────────── + +func (h *Handler) processServerAnnounce(data []byte) { + if len(data) < 12 { + return + } + h.versionMajor = binary.LittleEndian.Uint16(data[4:]) + h.versionMinor = binary.LittleEndian.Uint16(data[6:]) + h.clientID = binary.LittleEndian.Uint32(data[8:]) + slog.Debug("rdpdr: server announce", "major", h.versionMajor, + "minor", h.versionMinor, "clientID", h.clientID) + + // 客户端响应版本取 min(自身上限, 服务端)(FreeRDP rdpdr_main.c 同款): + // major 上限 1,minor 上限 0x000D。旧实现回 1.5(RDP5.1 时代语义)与 + // GENERAL 能力集布局错位,均已按 FreeRDP 源码修正。 + h.versionMajor = min16(1, h.versionMajor) + h.versionMinor = min16(RDPDR_VERSION_MINOR_RDP10X, h.versionMinor) + + // Client Announce Reply:VersionMajor(2) VersionMinor(2) ClientId(4) + b := make([]byte, 12) + binary.LittleEndian.PutUint16(b[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(b[2:], PAKID_CORE_CLIENTID_CONFIRM) + binary.LittleEndian.PutUint16(b[4:], h.versionMajor) + binary.LittleEndian.PutUint16(b[6:], h.versionMinor) + binary.LittleEndian.PutUint32(b[8:], h.clientID) + h.send(b) + + // Client Name Request:UnicodeFlag(4)=1 CodePage(4)=0 + // ComputerNameLen(4,含 NUL) ComputerName(UTF-16LE, NUL 结尾) + // 机器名取值依据见 clientComputerName 注释(M1 已知良好值 "tsclient"; + // "grdpclient" 变体失败与 CLOSE 探测错误混在一起,未分离归因)。 + uname := append(utf16Bytes(clientComputerName), 0, 0) + b2 := make([]byte, 16+len(uname)) + binary.LittleEndian.PutUint16(b2[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(b2[2:], PAKID_CORE_CLIENT_NAME) + binary.LittleEndian.PutUint32(b2[4:], 1) // UnicodeFlag + binary.LittleEndian.PutUint32(b2[8:], 0) // CodePage + binary.LittleEndian.PutUint32(b2[12:], uint32(len(uname))) + copy(b2[16:], uname) + h.send(b2) + + // 设备列表在 USER_LOGGEDON(且 CLIENTID_CONFIRM 已到)后发 + // (MS-RDPEFS 3.2.5.1.3;提前宣告服务端报 0xC0000001 拒装设备) +} + +func (h *Handler) sendDeviceList() { + h.deviceListSent = true + // DEVICE_ANNOUNCE{DeviceType(4) DeviceId(4) DosName(8) DeviceDataLength(4) + // DeviceData},与 FreeRDP drive_main.c/rdpdr_main.c 逐字节对齐: + // DosName = 共享名前 8 字节(短则 NUL 填充,高位字节替换 '_'); + // DeviceData = ASCII 全名 + 1 字节 NUL(FreeRDP 同款,Win10 接受; + // V02 协商下 DeviceDataLength 为 0 会遭服务端 0xC0000001 拒装)。 + dos := []byte(h.devName) + if len(dos) > 8 { + dos = dos[:8] + } + for i := range dos { + if dos[i] > 0x7F { + dos[i] = '_' + } + } + var dosName [8]byte + copy(dosName[:], dos) + devData := append([]byte(h.devName), 0) + // 头 4 + DeviceCount 4 + DeviceType 4 + DeviceId 4 + DosName 8 + + // DeviceDataLength 4 = 28,随后 DeviceData + b := make([]byte, 28+len(devData)) + binary.LittleEndian.PutUint16(b[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(b[2:], PAKID_CORE_DEVICELIST_ANNOUNCE) + binary.LittleEndian.PutUint32(b[4:], 1) // DeviceCount + binary.LittleEndian.PutUint32(b[8:], RDPDR_DTYP_FILESYSTEM) + binary.LittleEndian.PutUint32(b[12:], h.deviceID) + copy(b[16:], dosName[:]) + binary.LittleEndian.PutUint32(b[24:], uint32(len(devData))) + copy(b[28:], devData) + h.send(b) + slog.Debug("rdpdr: device list announced", "name", h.devName, "deviceID", h.deviceID, + "devDataLen", len(devData)) +} + +// ── 能力协商 ───────────────────────────────────────────────────────────── + +func (h *Handler) processServerCapability(data []byte) { + // 解析服务端 GENERAL capset,取 ioCode1(能力响应须与之求交) + if len(data) >= 12 { + // Server Capability:头 8 + numCapabilities(2) + Padding(2),随后 capset 列表 + off := 8 + for off+8 <= len(data) { + typ := binary.LittleEndian.Uint16(data[off:]) + length := int(binary.LittleEndian.Uint16(data[off+2:])) + if length < 8 || off+length > len(data) { + break + } + if typ == CAP_GENERAL_TYPE && length >= 44 { + // capset 内:osType(4) osVersion(4) protoMajor(2) protoMinor(2) + // ioCode1(4) ioCode2(4) extendedPDU(4) ... + h.serverIOCode1 = binary.LittleEndian.Uint32(data[off+20:]) + slog.Debug("rdpdr: server caps", "num", binary.LittleEndian.Uint16(data[4:]), + "ioCode1", h.serverIOCode1, + "extendedPDU", binary.LittleEndian.Uint32(data[off+28:])) + } + off += length + } + } + + // Client Core Capability Response(FreeRDP rdpdr_capabilities.c 逐字节对齐): + // 头 8 + GENERAL 44 + DRIVE 8 = 60。 + // 旧实现三处错误会使服务端拒建共享映射:GENERAL 的协议版本误写 4 字节 + //(服务端解析成 major=5 minor=0 ioCode1=0)、DRIVE capset 声明 10 写 12、 + // ioCode1=0/extendedPDU=4/extraFlags1=0。 + b := make([]byte, 8+44+8) + binary.LittleEndian.PutUint16(b[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(b[2:], PAKID_CORE_CLIENT_CAPABILITY) + binary.LittleEndian.PutUint16(b[4:], 2) // numCapabilities + binary.LittleEndian.PutUint16(b[6:], 0) // Padding + + // GENERAL:header(2+2+4) + osType(4) osVersion(4) protoMajor(2) protoMinor(2) + // ioCode1(4) ioCode2(4) extendedPDU(4) extraFlags1(4) extraFlags2(4) + // specialTypeDeviceCap(4) = 44 + binary.LittleEndian.PutUint16(b[8:], CAP_GENERAL_TYPE) + binary.LittleEndian.PutUint16(b[10:], 44) + binary.LittleEndian.PutUint32(b[12:], 2) // Version = GENERAL_CAPABILITY_VERSION_02 + // OsType @16 = 0, OsVersion @20 = 0 + binary.LittleEndian.PutUint16(b[24:], h.versionMajor) + binary.LittleEndian.PutUint16(b[26:], h.versionMinor) + binary.LittleEndian.PutUint32(b[28:], clientIOCode1&h.serverIOCode1) + // IoCode2 @32 = 0 + binary.LittleEndian.PutUint32(b[36:], RDPDR_DEVICE_REMOVE_PDUS| + RDPDR_CLIENT_DISPLAY_NAME_PDU|RDPDR_USER_LOGGEDON_PDU) + binary.LittleEndian.PutUint32(b[40:], RDPDR_ENABLE_ASYNCIO) + // ExtraFlags2 @44 = 0, SpecialTypeDeviceCap @48 = 0 + + // DRIVE:仅 8 字节头 {type, CapabilityLength=8, Version=2},无额外字段 + binary.LittleEndian.PutUint16(b[52:], CAP_DRIVE_TYPE) + binary.LittleEndian.PutUint16(b[54:], 8) + binary.LittleEndian.PutUint32(b[56:], 2) + + h.send(b) + slog.Debug("rdpdr: client caps sent") +} + +// ── Device I/O Request 分发 ────────────────────────────────────────────── + +func (h *Handler) processIORequest(data []byte) { + if len(data) < 24 { + return + } + deviceID := binary.LittleEndian.Uint32(data[4:]) + fileID := binary.LittleEndian.Uint32(data[8:]) + completionID := binary.LittleEndian.Uint32(data[12:]) + major := binary.LittleEndian.Uint32(data[16:]) + minor := binary.LittleEndian.Uint32(data[20:]) + if deviceID != h.deviceID { + slog.Warn("rdpdr: io for unknown device", "deviceID", deviceID) + return + } + p := &pending{deviceID: deviceID, fileID: fileID, major: major} + slog.Debug("rdpdr: io", "fileID", fileID, "completion", completionID, + "major", major, "minor", minor, "len", len(data)) + + switch major { + case IRP_MJ_CREATE: + // 头 24 + DesiredAccess(4) AllocationSize(8) FileAttributes(4) + // SharedAccess(4) CreateDisposition(4) CreateOptions(4) PathLength(4) = 56 + if len(data) < 56 { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + pathLen := int(binary.LittleEndian.Uint32(data[52:])) + path := "" + if pathLen > 0 && 56+pathLen <= len(data) { + path = utf16ToString(data[56 : 56+pathLen]) + } + // FileId 由客户端在 CREATE 时分配(服务端请求里的 FileId 为 0), + // 后续所有 IO 携带该号,桥接侧以此为句柄键。 + p.fileID = h.allocFileID() + h.track(completionID, p) + if h.fs != nil { + h.fs.Open(completionID, p.fileID, path) + } else { + h.takePending(completionID) + h.completeStatus(completionID, p, STATUS_DEVICE_NOT_READY, nil) + } + case IRP_MJ_READ: + if len(data) < 24+4+8 { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + length := binary.LittleEndian.Uint32(data[24:]) + offset := binary.LittleEndian.Uint64(data[28:]) + h.track(completionID, p) + if h.fs != nil { + h.fs.Read(completionID, fileID, offset, length) + } + case IRP_MJ_CLOSE: + h.track(completionID, p) + if h.fs != nil { + h.fs.Close(completionID, fileID) + } + case IRP_MJ_DIRECTORY_CONTROL: + if minor == IRP_MN_QUERY_DIRECTORY && len(data) >= 56 { + p.info = binary.LittleEndian.Uint32(data[24:]) + initial := data[28] != 0 + if initial { + h.enumMu.Lock() + h.enumPos[fileID] = 0 + h.enumMu.Unlock() + } + h.track(completionID, p) + if h.fs != nil { + h.fs.List(completionID, fileID, initial) + } + } else { + // 变更通知:声明不支持,服务端退化为轮询 + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + } + case IRP_MJ_QUERY_VOLUME_INFORMATION: + if len(data) < 24+4 { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + p.info = binary.LittleEndian.Uint32(data[24:]) + h.track(completionID, p) + if h.fs != nil { + h.fs.Volume(completionID, p.info) + } + case IRP_MJ_QUERY_INFORMATION: + if len(data) < 24+4 { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + p.info = binary.LittleEndian.Uint32(data[24:]) + h.track(completionID, p) + if h.fs != nil { + h.fs.Stat(completionID, fileID) + } + case IRP_MJ_WRITE, IRP_MJ_SET_INFORMATION: + // 只读阶段(M1):明确拒绝写路径 + h.completeStatus(completionID, p, STATUS_ACCESS_DENIED, nil) + case IRP_MJ_DEVICE_CONTROL, IRP_MJ_LOCK_CONTROL: + // FreeRDP 的 Discard 语义:未实现的 FSCTL/锁请求以 SUCCESS 空数据 + // 完成——NOT_IMPLEMENTED 会让服务端把整个设备标记为"不支持"。 + h.completeStatus(completionID, p, STATUS_SUCCESS, nil) + default: + slog.Debug("rdpdr: unhandled IRP", "major", major) + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + } +} + +func (h *Handler) track(completionID uint32, p *pending) { + h.pendingMu.Lock() + h.pending[completionID] = p + h.pendingMu.Unlock() +} + +func (h *Handler) takePending(completionID uint32) *pending { + h.pendingMu.Lock() + p := h.pending[completionID] + delete(h.pending, completionID) + h.pendingMu.Unlock() + return p +} + +// ── 完成回调(由异步桥在 wasm 侧调用)──────────────────────────────────── + +func (h *Handler) completeStatus(completionID uint32, p *pending, status uint32, extra []byte) { + if p == nil { + p = &pending{} + } + if p.major == IRP_MJ_CLOSE { + h.enumMu.Lock() + delete(h.enumPos, p.fileID) + h.enumMu.Unlock() + // Device Close Response 带固定 5 字节零 Padding(MS-RDPEFS + // 2.2.1.4.4;FreeRDP drive_process_irp_close 同款 Stream_Zero(5))。 + // 缺 padding 的 16 字节短响应虽在浏览期被服务端容忍,但为与服务端 + // 探测解析器逐字节一致(bb998f3 只修正了状态码,长度仍与 FreeRDP + // 不同),这里统一补齐。 + extra = append(extra, 0, 0, 0, 0, 0) + } + slog.Debug("rdpdr: complete", "completion", completionID, "major", p.major, + "info", p.info, "status", status, "extra", len(extra)) + b := make([]byte, 16, 16+len(extra)) + binary.LittleEndian.PutUint16(b[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(b[2:], PAKID_CORE_DEVICE_IOCOMPLETION) + binary.LittleEndian.PutUint32(b[4:], p.deviceID) + binary.LittleEndian.PutUint32(b[8:], completionID) + binary.LittleEndian.PutUint32(b[12:], status) + b = append(b, extra...) + h.send(b) +} + +// CompleteStatus 完成一个纯状态响应(CLOSE/错误路径等)。 +func (h *Handler) CompleteStatus(completionID uint32, status uint32) { + h.completeStatus(completionID, h.takePending(completionID), status, nil) +} + +// CompleteBytes 完成带数据块的响应(READ:Length(4) 前缀 + 数据)。 +func (h *Handler) CompleteBytes(completionID uint32, status uint32, data []byte) { + p := h.takePending(completionID) + extra := make([]byte, 4+len(data)) + binary.LittleEndian.PutUint32(extra, uint32(len(data))) + copy(extra[4:], data) + h.completeStatus(completionID, p, status, extra) +} + +// CompleteJSON 完成结构化响应;kind 决定编码: +// - "create" json="dir"|"file":CREATE 响应(FileId+Information(FILE_OPENED)) +// - "list" json=[{name,dir,size,mtime,created,accessed}...]:按信息类编码目录条目 +// - "stat" json={size,mtime,created,accessed,dir}:QUERY_INFORMATION 响应 +// - "volume" json 忽略:QUERY_VOLUME_INFORMATION 响应(卷标取 Handler 配置) +// - "status" json 忽略:纯状态响应(CLOSE 等无载荷完成)——CLOSE 必须 +// 以 STATUS_SUCCESS 完成,返回错误状态(如 NOT_IMPLEMENTED)会让服务端 +// 判定设备异常、撤销 \\tsclient\<名> 映射(表现为"试图访问无效的地址" +// 且后续零 IRP,2026-09-12 10.0.0.3 服务端重启后探测序列新增 CLOSE 步骤 +// 时暴露)。 +// +// status != STATUS_SUCCESS 时一律回纯状态响应。 +func (h *Handler) CompleteJSON(completionID uint32, status uint32, kind, json string) { + p := h.takePending(completionID) + if p == nil { + // 未知完成号(重复回调/通道重置):回空 pending 的纯状态响应兜底 + h.completeStatus(completionID, &pending{}, status, nil) + return + } + if status != STATUS_SUCCESS { + h.completeStatus(completionID, p, status, nil) + return + } + switch kind { + case "status": + h.completeStatus(completionID, p, STATUS_SUCCESS, nil) + case "create": + h.enumMu.Lock() + h.enumPos[p.fileID] = 0 + h.enumMu.Unlock() + extra := make([]byte, 5) + binary.LittleEndian.PutUint32(extra, p.fileID) + extra[4] = 1 // Information = FILE_OPENED + h.completeStatus(completionID, p, STATUS_SUCCESS, extra) + case "list": + entries, err := decodeEntries(json) + if err != nil { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + h.enumMu.Lock() + pos := h.enumPos[p.fileID] + if pos > len(entries) || pos < 0 { + pos = 0 // 目录内容变化的兜底 + } + if pos >= len(entries) { + h.enumPos[p.fileID] = 0 + h.enumMu.Unlock() + h.completeStatus(completionID, p, STATUS_NO_MORE_FILES, nil) + return + } + batch := entries[pos:] + h.enumPos[p.fileID] = len(entries) + h.enumMu.Unlock() + data := encodeDirEntries(p.info, batch) + if data == nil { + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + return + } + extra := make([]byte, 4+len(data)) + binary.LittleEndian.PutUint32(extra, uint32(len(data))) + copy(extra[4:], data) + h.completeStatus(completionID, p, STATUS_SUCCESS, extra) + case "stat": + e, err := decodeEntry(json) + if err != nil { + h.completeStatus(completionID, p, STATUS_INVALID_PARAMETER, nil) + return + } + data := encodeFileInfo(p.info, e) + if data == nil { + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + return + } + extra := make([]byte, 4+len(data)) + binary.LittleEndian.PutUint32(extra, uint32(len(data))) + copy(extra[4:], data) + h.completeStatus(completionID, p, STATUS_SUCCESS, extra) + case "volume": + data := encodeVolumeInfo(p.info, h.label) + if data == nil { + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + return + } + extra := make([]byte, 4+len(data)) + binary.LittleEndian.PutUint32(extra, uint32(len(data))) + copy(extra[4:], data) + h.completeStatus(completionID, p, STATUS_SUCCESS, extra) + default: + h.completeStatus(completionID, p, STATUS_NOT_IMPLEMENTED, nil) + } +} + +// allocFileID 为新 CREATE 分配 FileId(服务端请求中 FileId=0,响应才带号)。 +func (h *Handler) allocFileID() uint32 { + id := h.nextFileID + h.nextFileID++ + return id +} diff --git a/plugin/rdpdr/rdpdr_test.go b/plugin/rdpdr/rdpdr_test.go new file mode 100644 index 0000000..4181606 --- /dev/null +++ b/plugin/rdpdr/rdpdr_test.go @@ -0,0 +1,428 @@ +package rdpdr + +import ( + "bytes" + "encoding/binary" + "testing" +) + +// fakeSender 捕获发往通道的消息。 +type fakeSender struct{ msgs [][]byte } + +func (f *fakeSender) SendToChannel(ch string, s []byte) (int, error) { + f.msgs = append(f.msgs, append([]byte(nil), s...)) + return len(s), nil +} + +// fakeFS 记录派发并支持手动完成。 +type fakeFS struct { + h *Handler + opens []string + lists int + complet []uint32 +} + +func (f *fakeFS) Open(completionID, fileID uint32, path string) { + f.opens = append(f.opens, path) +} +func (f *fakeFS) Read(completionID, fileID uint32, offset uint64, length uint32) {} +func (f *fakeFS) Close(completionID, fileID uint32) {} +func (f *fakeFS) List(completionID, fileID uint32, initial bool) { + f.complet = append(f.complet, completionID) + f.lists++ +} +func (f *fakeFS) Stat(completionID, fileID uint32) {} +func (f *fakeFS) Volume(completionID uint32, infoClass uint32) {} + +func newTestHandler() (*Handler, *fakeSender, *fakeFS) { + h := NewHandler("webrdp", "local") + s := &fakeSender{} + h.Sender(s) + fs := &fakeFS{h: h} + h.SetFilesystem(fs) + return h, s, fs +} + +func TestServerAnnounceFlow(t *testing.T) { + h, s, _ := newTestHandler() + // Server Announce:major 1 minor 5 clientID 0x1234 + req := make([]byte, 12) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_SERVER_ANNOUNCE) + binary.LittleEndian.PutUint16(req[4:], 1) + binary.LittleEndian.PutUint16(req[6:], 5) + binary.LittleEndian.PutUint32(req[8:], 0x1234) + h.Process(req) + + if len(s.msgs) != 2 { + t.Fatalf("应发 2 条消息(announce reply/name),实发 %d", len(s.msgs)) + } + // 1) Client Announce Reply:v1.5 + 原 ClientId + if got := binary.LittleEndian.Uint16(s.msgs[0][2:]); got != PAKID_CORE_CLIENTID_CONFIRM { + t.Fatalf("msg0 packetID=%#x", got) + } + if binary.LittleEndian.Uint16(s.msgs[0][4:]) != 1 || binary.LittleEndian.Uint16(s.msgs[0][6:]) != 5 { + t.Fatal("announce reply 版本应为 1.5") + } + if binary.LittleEndian.Uint32(s.msgs[0][8:]) != 0x1234 { + t.Fatal("announce reply 应原样回 ClientId") + } + // 2) Client Name Request:UnicodeFlag=1 + if got := binary.LittleEndian.Uint16(s.msgs[1][2:]); got != PAKID_CORE_CLIENT_NAME { + t.Fatalf("msg1 packetID=%#x", got) + } + if binary.LittleEndian.Uint32(s.msgs[1][4:]) != 1 { + t.Fatal("UnicodeFlag 应为 1") + } + + // 服务端 CLIENTID_CONFIRM:仅记录状态,设备列表还不发 + req2 := make([]byte, 8) + binary.LittleEndian.PutUint16(req2[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req2[2:], PAKID_CORE_CLIENTID_CONFIRM) + h.Process(req2) + if len(s.msgs) != 2 { + t.Fatalf("CLIENTID_CONFIRM 后不应发消息,实发 %d", len(s.msgs)) + } + // 服务端 USER_LOGGEDON → 此刻才宣告设备列表(MS-RDPEFS 3.2.5.1.3) + req3 := make([]byte, 8) + binary.LittleEndian.PutUint16(req3[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req3[2:], PAKID_CORE_USER_LOGGEDON) + h.Process(req3) + if len(s.msgs) != 3 { + t.Fatalf("USER_LOGGEDON 后应发设备列表,实发 %d", len(s.msgs)) + } + // 3) Device List:count=1, type=FILESYSTEM, DosName "webrdp" + if got := binary.LittleEndian.Uint16(s.msgs[2][2:]); got != PAKID_CORE_DEVICELIST_ANNOUNCE { + t.Fatalf("msg2 packetID=%#x", got) + } + if binary.LittleEndian.Uint32(s.msgs[2][4:]) != 1 { + t.Fatal("DeviceCount 应为 1") + } + if binary.LittleEndian.Uint32(s.msgs[2][8:]) != RDPDR_DTYP_FILESYSTEM { + t.Fatal("DeviceType 应为 FILESYSTEM(8)") + } + if got := binary.LittleEndian.Uint32(s.msgs[2][12:]); got != 1 { + t.Fatalf("DeviceId 应为 1(非 0),实 %d", got) + } + if name := string(bytes.TrimRight(s.msgs[2][16:24], "\x00")); name != "webrdp" { + t.Fatalf("DosName=%q", name) + } + // V02 协商下 DeviceData = ASCII 全名 + 1 字节 NUL(FreeRDP drive_main.c + // 同款;服务端 UNC 映射键取 DosName,DeviceData 内容须非空否则拒装) + if got := len(s.msgs[2]); got != 35 { // 28 + "webrdp" ASCII 6 + NUL 1 + t.Fatalf("设备列表报文应为 35 字节,实 %d", got) + } + if got := binary.LittleEndian.Uint32(s.msgs[2][24:]); got != 7 { + t.Fatalf("DeviceDataLength 应为 7,实 %d", got) + } + if !bytes.Equal(s.msgs[2][28:35], []byte("webrdp\x00")) { + t.Fatalf("DeviceData 应为 ASCII 全名+NUL,实 % x", s.msgs[2][28:35]) + } +} + +func TestAnnounceVersionNegotiation(t *testing.T) { + h, s, _ := newTestHandler() + // 服务端 1.13(Win10)→ 客户端应回 1.MIN(0x000D, 0x000D)=1.13 + req := make([]byte, 12) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_SERVER_ANNOUNCE) + binary.LittleEndian.PutUint16(req[4:], 1) + binary.LittleEndian.PutUint16(req[6:], 0x000D) + binary.LittleEndian.PutUint32(req[8:], 0x1234) + h.Process(req) + if binary.LittleEndian.Uint16(s.msgs[0][4:]) != 1 || + binary.LittleEndian.Uint16(s.msgs[0][6:]) != 0x000D { + t.Fatalf("服务端 1.13 时应回 1.13,实 %d.%d", + binary.LittleEndian.Uint16(s.msgs[0][4:]), + binary.LittleEndian.Uint16(s.msgs[0][6:])) + } + // 服务端 1.5(老语义)→ 回 1.5 + h2, s2, _ := newTestHandler() + binary.LittleEndian.PutUint16(req[6:], 5) + h2.Process(req) + if binary.LittleEndian.Uint16(s2.msgs[0][6:]) != 5 { + t.Fatal("服务端 1.5 时应回 1.5") + } +} + +func TestSingleAnnounceAtLogon(t *testing.T) { + h, s, _ := newTestHandler() + // Server Announce → Server Caps → ClientID Confirm:1.0x 语义下 + // 均不应发设备列表(FreeRDP:非登录阶段跳过 FS 设备,count=0 不发) + ann := make([]byte, 12) + binary.LittleEndian.PutUint16(ann[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(ann[2:], PAKID_CORE_SERVER_ANNOUNCE) + binary.LittleEndian.PutUint16(ann[4:], 1) + binary.LittleEndian.PutUint16(ann[6:], 0x000D) + binary.LittleEndian.PutUint32(ann[8:], 0x1234) + h.Process(ann) + caps := make([]byte, 8) + binary.LittleEndian.PutUint16(caps[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(caps[2:], PAKID_CORE_SERVER_CAPABILITY) + caps[4] = 0 + h.Process(caps) + conf := make([]byte, 12) + binary.LittleEndian.PutUint16(conf[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(conf[2:], PAKID_CORE_CLIENTID_CONFIRM) + binary.LittleEndian.PutUint16(conf[4:], 1) + binary.LittleEndian.PutUint16(conf[6:], 0x000D) + binary.LittleEndian.PutUint32(conf[8:], 0x99) // 服务端回显采纳 + h.Process(conf) + // msgs: announce reply, name request, caps response——无设备列表 + if n := len(s.msgs); n != 3 { + t.Fatalf("登录前不应发设备列表(3 条),实 %d", n) + } + if h.clientID != 0x99 { + t.Fatalf("应采纳 Confirm 回显的 ClientId,实 %#x", h.clientID) + } + // USER_LOGGEDON:唯一一次设备列表宣告 + logon := make([]byte, 8) + binary.LittleEndian.PutUint16(logon[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(logon[2:], PAKID_CORE_USER_LOGGEDON) + h.Process(logon) + if n := len(s.msgs); n != 4 { + t.Fatalf("USER_LOGGEDON 后应发设备列表(4 条),实 %d", n) + } + if got := binary.LittleEndian.Uint16(s.msgs[3][2:]); got != PAKID_CORE_DEVICELIST_ANNOUNCE { + t.Fatalf("msg3 应为设备列表,实 %#x", got) + } + // 重复 USER_LOGGEDON 不重发 + h.Process(logon) + if n := len(s.msgs); n != 4 { + t.Fatalf("重复 USER_LOGGEDON 不应重发,实 %d", n) + } +} + +func TestSanitizeShareName(t *testing.T) { + if got := sanitizeShareName(`a:bd"e/f\g|h i,j`); got != "a_b_c_d_e_f_g_h_i_j" { + t.Fatalf("sanitizeShareName=%q", got) + } + if got := sanitizeShareName("rdpdrive-test"); got != "rdpdrive-test" { + t.Fatalf("合法名不应改动,实 %q", got) + } +} + +func TestClientCapabilityResponse(t *testing.T) { + h, s, _ := newTestHandler() + // 先握手(服务端 1.13 → 客户端版本 1.13),再发带 GENERAL capset 的 + // 服务端能力请求(ioCode1 = 全部常见 IRP 位) + ann := make([]byte, 12) + binary.LittleEndian.PutUint16(ann[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(ann[2:], PAKID_CORE_SERVER_ANNOUNCE) + binary.LittleEndian.PutUint16(ann[4:], 1) + binary.LittleEndian.PutUint16(ann[6:], 0x000D) + binary.LittleEndian.PutUint32(ann[8:], 0x1234) + h.Process(ann) + + req := make([]byte, 8+44) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_SERVER_CAPABILITY) + binary.LittleEndian.PutUint16(req[4:], 1) // numCapabilities + binary.LittleEndian.PutUint16(req[8:], CAP_GENERAL_TYPE) + binary.LittleEndian.PutUint16(req[10:], 44) + binary.LittleEndian.PutUint32(req[12:], 2) + binary.LittleEndian.PutUint32(req[28:], 0xFFFFFFFF) // ioCode1 + h.Process(req) + + if len(s.msgs) != 3 || len(s.msgs[2]) != 8+44+8 { + t.Fatalf("能力响应长度应 60,实 %v", len(s.msgs[2])) + } + m := s.msgs[2] + if binary.LittleEndian.Uint16(m[2:]) != PAKID_CORE_CLIENT_CAPABILITY { + t.Fatal("packetID 应为 CLIENT_CAPABILITY") + } + if binary.LittleEndian.Uint16(m[4:]) != 2 { + t.Fatal("numCapabilities 应为 2") + } + // GENERAL:版本字段各 2 字节(旧实现误写 4 字节,服务端解析错位) + if binary.LittleEndian.Uint16(m[8:]) != CAP_GENERAL_TYPE || binary.LittleEndian.Uint16(m[10:]) != 44 { + t.Fatal("GENERAL capset 头错误") + } + if binary.LittleEndian.Uint16(m[24:]) != 1 || binary.LittleEndian.Uint16(m[26:]) != 0x000D { + t.Fatal("GENERAL 协议版本应为 1.13(各 2 字节)") + } + // ioCode1 = 客户端掩码 ∩ 服务端掩码 + if got := binary.LittleEndian.Uint32(m[28:]); got != clientIOCode1 { + t.Fatalf("ioCode1 应为求交结果 %#x,实 %#x", clientIOCode1, got) + } + // extendedPDU = REMOVE|DISPLAY_NAME|USER_LOGGEDON;extraFlags1 = ENABLE_ASYNCIO + if got := binary.LittleEndian.Uint32(m[36:]); got != 7 { + t.Fatalf("extendedPDU 应为 7,实 %#x", got) + } + if got := binary.LittleEndian.Uint32(m[40:]); got != RDPDR_ENABLE_ASYNCIO { + t.Fatalf("extraFlags1 应为 ENABLE_ASYNCIO,实 %#x", got) + } + // DRIVE:仅 8 字节头 {type=4, len=8, version=2}(FreeRDP 同款) + if binary.LittleEndian.Uint16(m[52:]) != CAP_DRIVE_TYPE || binary.LittleEndian.Uint16(m[54:]) != 8 { + t.Fatal("DRIVE capset 头错误(FreeRDP 对齐:仅 8 字节头,len=8)") + } + if binary.LittleEndian.Uint32(m[56:]) != 2 { + t.Fatal("DRIVE capset Version 应为 2") + } +} + +func TestCreateDispatchAndComplete(t *testing.T) { + h, s, fs := newTestHandler() + + // 构造 CREATE 请求:path "\hello.txt"(FileId 由 handler 分配,预期为 1) + path := utf16Bytes("\\hello.txt") + req := make([]byte, 56+len(path)) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_DEVICE_IOREQUEST) + binary.LittleEndian.PutUint32(req[4:], h.deviceID) + binary.LittleEndian.PutUint32(req[12:], 77) // CompletionId + binary.LittleEndian.PutUint32(req[16:], IRP_MJ_CREATE) + binary.LittleEndian.PutUint32(req[52:], uint32(len(path))) + copy(req[56:], path) + h.Process(req) + + if len(fs.opens) != 1 || fs.opens[0] != "\\hello.txt" { + t.Fatalf("Open 未正确派发: %v", fs.opens) + } + + // 完成路径 A:失败 → 纯状态响应 + h.CompleteJSON(77, STATUS_NO_SUCH_FILE, "create", "file") + if len(s.msgs) != 1 { + t.Fatalf("失败完成应发 1 条响应,实 %d", len(s.msgs)) + } + resp := s.msgs[0] + if binary.LittleEndian.Uint16(resp[2:]) != PAKID_CORE_DEVICE_IOCOMPLETION || + binary.LittleEndian.Uint32(resp[8:]) != 77 || + binary.LittleEndian.Uint32(resp[12:]) != STATUS_NO_SUCH_FILE { + t.Fatal("失败 CREATE 响应头错误") + } + + // 完成路径 B:成功 → 16 字节头 + FileId(4) + Information(1) + h.track(78, &pending{deviceID: h.deviceID, fileID: 1, major: IRP_MJ_CREATE}) + h.CompleteJSON(78, STATUS_SUCCESS, "create", "file") + resp = s.msgs[1] + if len(resp) != 21 { + t.Fatalf("成功 CREATE 响应应 21 字节,实 %d", len(resp)) + } + if binary.LittleEndian.Uint32(resp[16:]) != 1 || resp[20] != 1 { + t.Fatal("CREATE 响应 FileId/Information 错误") + } +} + +func TestDirectoryListPagination(t *testing.T) { + h, s, fs := newTestHandler() + dirID := uint32(1) + + req := make([]byte, 56) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_DEVICE_IOREQUEST) + binary.LittleEndian.PutUint32(req[4:], h.deviceID) + binary.LittleEndian.PutUint32(req[8:], dirID) + binary.LittleEndian.PutUint32(req[12:], 500) + binary.LittleEndian.PutUint32(req[16:], IRP_MJ_DIRECTORY_CONTROL) + binary.LittleEndian.PutUint32(req[20:], IRP_MN_QUERY_DIRECTORY) + binary.LittleEndian.PutUint32(req[24:], FileBothDirectoryInformation) + req[28] = 1 // InitialQuery + h.Process(req) + + if len(fs.complet) != 1 { + t.Fatal("List 未派发") + } + // 桥接返回 2 条目 + h.CompleteJSON(fs.complet[0], STATUS_SUCCESS, "list", + `[{"name":"a.txt","size":5,"mtime":1700000000000},{"name":"sub","dir":true,"mtime":1700000000000}]`) + resp := s.msgs[0] + if binary.LittleEndian.Uint32(resp[12:]) != STATUS_SUCCESS { + t.Fatal("首次枚举应成功") + } + dataLen := binary.LittleEndian.Uint32(resp[16:]) + data := resp[20:] + if int(dataLen) != len(data)-0 || dataLen == 0 { + t.Fatalf("数据长度字段不一致: %d", dataLen) + } + // 条目链校验:第一条 NextEntryOffset 指向第二条,第二条为 0 + first := binary.LittleEndian.Uint32(data[0:]) + if first == 0 || int(first)+4 > len(data) { + t.Fatalf("NextEntryOffset 链错误: %d", first) + } + // a.txt 名字长度(UTF-16 字节数)在 class3: FileNameLength @60 + nameLen := binary.LittleEndian.Uint32(data[60:]) + if nameLen != 10 { // "a.txt" 5 chars × 2 + t.Fatalf("FileNameLength=%d", nameLen) + } + second := binary.LittleEndian.Uint32(data[first:]) + if second != 0 { + t.Fatalf("末条 NextEntryOffset 应为 0,实 %d", second) + } + + // 续枚举(InitialQuery=0):桥接再次返回全量,游标已到末尾 → NO_MORE_FILES + req[28] = 0 + h.Process(req) + if len(fs.complet) != 2 { + t.Fatal("续枚举 List 未派发") + } + h.CompleteJSON(fs.complet[1], STATUS_SUCCESS, "list", + `[{"name":"a.txt","size":5,"mtime":1700000000000},{"name":"sub","dir":true,"mtime":1700000000000}]`) + resp = s.msgs[1] + if binary.LittleEndian.Uint32(resp[12:]) != STATUS_NO_MORE_FILES { + t.Fatalf("续枚举应返回 NO_MORE_FILES,实 %#x", binary.LittleEndian.Uint32(resp[12:])) + } +} + +func TestReadCompletionEncoding(t *testing.T) { + h, s, _ := newTestHandler() + h.track(9, &pending{deviceID: h.deviceID, fileID: 3, major: IRP_MJ_READ}) + h.CompleteBytes(9, STATUS_SUCCESS, []byte("hello")) + resp := s.msgs[0] + if len(resp) != 16+4+5 { + t.Fatalf("READ 响应应 25 字节,实 %d", len(resp)) + } + if binary.LittleEndian.Uint32(resp[16:]) != 5 { + t.Fatal("Length 前缀应为 5") + } + if !bytes.Equal(resp[20:], []byte("hello")) { + t.Fatal("数据内容不一致") + } +} + +func TestWriteDeniedInReadOnlyMode(t *testing.T) { + h, s, _ := newTestHandler() + req := make([]byte, 24) + binary.LittleEndian.PutUint16(req[0:], RDPDR_CTYP_CORE) + binary.LittleEndian.PutUint16(req[2:], PAKID_CORE_DEVICE_IOREQUEST) + binary.LittleEndian.PutUint32(req[4:], h.deviceID) + binary.LittleEndian.PutUint32(req[12:], 42) + binary.LittleEndian.PutUint32(req[16:], IRP_MJ_WRITE) + h.Process(req) + if len(s.msgs) != 1 { + t.Fatal("WRITE 应直接拒绝") + } + if binary.LittleEndian.Uint32(s.msgs[0][12:]) != STATUS_ACCESS_DENIED { + t.Fatal("WRITE 应回 ACCESS_DENIED") + } +} + +func TestDirEntryEncodingLayouts(t *testing.T) { + entries := []DirEntry{{Name: "x", Dir: true, Size: 0, Mtime: 1700000000000}} + // 单条即末条:NextEntryOffset=0,无对齐填充 + if d := encodeDirEntries(FileDirectoryInformation, entries); len(d) != 64+2 { + t.Fatalf("class1 单条应为 66,实 %d", len(d)) + } + if d := encodeDirEntries(FileBothDirectoryInformation, entries); len(d) != 93+2 { + t.Fatalf("class3 单条应为 95,实 %d", len(d)) + } + if d := encodeDirEntries(FileNamesInformation, entries); len(d) != 12+2 { + t.Fatalf("names 单条应为 14,实 %d", len(d)) + } + if d := encodeDirEntries(0x99, entries); d != nil { + t.Fatal("未知信息类应返回 nil") + } + // 非末条需 4 字节对齐:两条 class1("x"=2 字节名)→ 68 + 66 + two := []DirEntry{{Name: "x", Mtime: 1}, {Name: "y", Mtime: 1}} + if d := encodeDirEntries(FileDirectoryInformation, two); len(d) != 68+66 { + t.Fatalf("class1 两条应 68+66,实 %d", len(d)) + } + // FileBasicInformation 40 字节 + if d := encodeFileInfo(FileBasicInformation, entries[0]); len(d) != 40 { + t.Fatalf("FileBasicInformation 应 40 字节,实 %d", len(d)) + } + // FileFsSizeInformation 24 字节 + if d := encodeVolumeInfo(FileFsSizeInformation, "local"); len(d) != 24 { + t.Fatalf("FileFsSizeInformation 应 24 字节,实 %d", len(d)) + } +} diff --git a/plugin/rdpedisp/rdpedisp.go b/plugin/rdpedisp/rdpedisp.go new file mode 100644 index 0000000..ebeadd7 --- /dev/null +++ b/plugin/rdpedisp/rdpedisp.go @@ -0,0 +1,186 @@ +// Package rdpedisp implements the RDP Display Update Virtual Channel +// (MS-RDPEDISP), which allows the client to request a resolution change +// while connected. The channel name is: +// +// "Microsoft::Windows::RDS::DisplayControl" +// +// Typical usage: +// +// 1. Register the handler with the DVC client before connecting. +// 2. After the session is established, call SendMonitorLayout to request +// a new resolution. The server will reshape the desktop and send a +// fresh RDPGFX ResetGraphics command. +package rdpedisp + +import ( + "encoding/binary" + "log/slog" + "sync" +) + +// ChannelName is the well-known DVC name for the Display Update channel. +const ChannelName = "Microsoft::Windows::RDS::DisplayControl" + +// PDU types (MS-RDPEDISP 2.2.1.2.1) +const ( + pduTypeCaps = 0x00000005 + pduTypeMonitorLayout = 0x00000002 +) + +// DISPLAYCONTROL_MONITOR_PRIMARY marks the primary monitor. +const MonitorFlagPrimary = uint32(0x00000001) + +// monitorLayoutSize is the fixed per-monitor record size required by the spec. +const monitorLayoutSize = 40 + +// Monitor describes a single monitor in a MONITOR_LAYOUT PDU. +// Width and Height must each be at least 200 and Width must be even. +// Set PhysicalWidth/Height to 0 when the physical dimensions are unknown. +// Orientation: 0=landscape (normal), 90, 180, 270. +// DesktopScaleFactor: one of 100, 125, 150, 175, 200 (use 100 if unsure). +// DeviceScaleFactor: one of 100, 140, 180 (use 100 if unsure). +type Monitor struct { + Flags uint32 + Left int32 + Top int32 + Width uint32 + Height uint32 + PhysicalWidth uint32 + PhysicalHeight uint32 + Orientation uint32 + DesktopScaleFactor uint32 + DeviceScaleFactor uint32 +} + +// Handler is the DVC handler for the Display Update channel. +// It implements the drdynvc.DvcChannelHandler interface and the optional +// SetSendFunc / OnChannelCreated extension interfaces. +type Handler struct { + mu sync.Mutex + send func([]byte) + capsReceived bool + maxMonitors uint32 + pendingMonitors []Monitor +} + +// NewHandler returns a new Handler. +// width and height are the desired initial desktop dimensions; when both are +// non-zero the handler queues an initial layout change that is sent as soon as +// the server advertises its capabilities (MS-RDPEDISP CAPS PDU), prompting +// servers such as GNOME Remote Desktop (headless) to resize. +func NewHandler(width, height uint32) *Handler { + h := &Handler{} + if width > 0 && height > 0 { + h.pendingMonitors = []Monitor{ + { + Flags: MonitorFlagPrimary, + Left: 0, + Top: 0, + Width: width, + Height: height, + PhysicalWidth: 0, + PhysicalHeight: 0, + Orientation: 0, + DesktopScaleFactor: 100, + DeviceScaleFactor: 100, + }, + } + } + return h +} + +// SetSendFunc is called by the DVC client to provide a write-back function. +// Required by the drdynvc channel plumbing. +func (h *Handler) SetSendFunc(f func([]byte)) { + h.mu.Lock() + defer h.mu.Unlock() + h.send = f +} + +// OnChannelCreated is called by the DVC client after the CREATE_RSP has been sent. +func (h *Handler) OnChannelCreated() { + slog.Debug("rdpedisp: channel created") +} + +// Process handles incoming data from the server (CAPS PDU, etc.). +func (h *Handler) Process(data []byte) { + if len(data) < 8 { + return + } + pduType := binary.LittleEndian.Uint32(data[0:4]) + switch pduType { + case pduTypeCaps: + if len(data) >= 20 { + maxMonitors := binary.LittleEndian.Uint32(data[8:12]) + slog.Debug("rdpedisp: server CAPS", "maxMonitors", maxMonitors) + h.mu.Lock() + h.capsReceived = true + h.maxMonitors = maxMonitors + pending := h.pendingMonitors + h.pendingMonitors = nil + h.mu.Unlock() + + if len(pending) > 0 { + h.SendMonitorLayout(pending) + } + } + default: + slog.Debug("rdpedisp: unknown PDU type", "type", pduType) + } +} + +// SendMonitorLayout sends a DISPLAYCONTROL_MONITOR_LAYOUT_PDU to the server, +// requesting the given monitor configuration. Per MS-RDPEDISP 3.2.5.1, the +// client MUST NOT send a layout PDU before receiving the server CAPS PDU; if +// called before CAPS is received, the layout is queued and sent automatically +// when CAPS arrives. +// +// The server will apply the new layout and—if using the RDPGFX pipeline—send +// a ResetGraphics command that resets surface dimensions to match. +func (h *Handler) SendMonitorLayout(monitors []Monitor) { + h.mu.Lock() + if !h.capsReceived { + slog.Debug("rdpedisp: queuing MonitorLayout (waiting for server CAPS)", "numMonitors", len(monitors)) + h.pendingMonitors = monitors + h.mu.Unlock() + return + } + send := h.send + h.mu.Unlock() + + if send == nil { + slog.Warn("rdpedisp: SendMonitorLayout: channel not open") + return + } + + numMonitors := uint32(len(monitors)) + // PDU layout: 8-byte DISPLAYCONTROL_HEADER + + // 4 bytes MonitorLayoutSize + + // 4 bytes NumMonitors + + // numMonitors * monitorLayoutSize bytes + pduLen := uint32(8 + 4 + 4 + int(numMonitors)*monitorLayoutSize) + pdu := make([]byte, pduLen) + + binary.LittleEndian.PutUint32(pdu[0:], pduTypeMonitorLayout) + binary.LittleEndian.PutUint32(pdu[4:], pduLen) + binary.LittleEndian.PutUint32(pdu[8:], monitorLayoutSize) // fixed record size + binary.LittleEndian.PutUint32(pdu[12:], numMonitors) + + off := 16 + for _, m := range monitors { + binary.LittleEndian.PutUint32(pdu[off+0:], m.Flags) + binary.LittleEndian.PutUint32(pdu[off+4:], uint32(m.Left)) + binary.LittleEndian.PutUint32(pdu[off+8:], uint32(m.Top)) + binary.LittleEndian.PutUint32(pdu[off+12:], m.Width) + binary.LittleEndian.PutUint32(pdu[off+16:], m.Height) + binary.LittleEndian.PutUint32(pdu[off+20:], m.PhysicalWidth) + binary.LittleEndian.PutUint32(pdu[off+24:], m.PhysicalHeight) + binary.LittleEndian.PutUint32(pdu[off+28:], m.Orientation) + binary.LittleEndian.PutUint32(pdu[off+32:], m.DesktopScaleFactor) + binary.LittleEndian.PutUint32(pdu[off+36:], m.DeviceScaleFactor) + off += monitorLayoutSize + } + + slog.Debug("rdpedisp: sending MonitorLayout", "numMonitors", numMonitors) + send(pdu) +} diff --git a/plugin/rdpedisp/rdpedisp_test.go b/plugin/rdpedisp/rdpedisp_test.go new file mode 100644 index 0000000..834f89b --- /dev/null +++ b/plugin/rdpedisp/rdpedisp_test.go @@ -0,0 +1,63 @@ +package rdpedisp + +import ( + "encoding/binary" + "testing" +) + +func TestRdpedispQueuesBeforeCaps(t *testing.T) { + h := NewHandler(1920, 1080) + var sent [][]byte + h.SetSendFunc(func(data []byte) { + copied := make([]byte, len(data)) + copy(copied, data) + sent = append(sent, copied) + }) + + // OnChannelCreated should NOT send anything + h.OnChannelCreated() + if len(sent) != 0 { + t.Fatalf("expected 0 packets before CAPS, got %d", len(sent)) + } + + // Server sends CAPS PDU + capsData := make([]byte, 20) + binary.LittleEndian.PutUint32(capsData[0:], pduTypeCaps) + binary.LittleEndian.PutUint32(capsData[4:], 20) + binary.LittleEndian.PutUint32(capsData[8:], 16) // MaxNumMonitors = 16 + binary.LittleEndian.PutUint32(capsData[12:], 4096) + binary.LittleEndian.PutUint32(capsData[16:], 2048) + + h.Process(capsData) + + if len(sent) != 1 { + t.Fatalf("expected 1 queued packet sent after CAPS, got %d", len(sent)) + } + + pdu := sent[0] + if len(pdu) < 16+monitorLayoutSize { + t.Fatalf("invalid pdu length: %d", len(pdu)) + } + pduType := binary.LittleEndian.Uint32(pdu[0:]) + if pduType != pduTypeMonitorLayout { + t.Fatalf("expected pduTypeMonitorLayout (2), got %d", pduType) + } + w := binary.LittleEndian.Uint32(pdu[16+12:]) + hVal := binary.LittleEndian.Uint32(pdu[16+16:]) + if w != 1920 || hVal != 1080 { + t.Fatalf("expected 1920x1080, got %dx%d", w, hVal) + } + + // Subsequent SendMonitorLayout sends immediately + h.SendMonitorLayout([]Monitor{ + { + Flags: MonitorFlagPrimary, + Width: 1280, + Height: 720, + }, + }) + + if len(sent) != 2 { + t.Fatalf("expected 2 packets, got %d", len(sent)) + } +} diff --git a/plugin/rdpgfx/avc.go b/plugin/rdpgfx/avc.go new file mode 100644 index 0000000..deb1cf8 --- /dev/null +++ b/plugin/rdpgfx/avc.go @@ -0,0 +1,2329 @@ +package rdpgfx + +// AVC420 / AVC444 bitmap stream parsing (MS-RDPEGFX 2.2.4.6 / 2.2.4.7). + +import ( + "encoding/binary" + "fmt" + "log/slog" + "runtime" + "sync" + "time" +) + +type avcRect struct { + left, top, right, bottom uint16 +} + +type avc420Stream struct { + regions []avcRect + h264Data []byte +} + +// fillAVC420Stream parses data into out in-place, reusing out.regions if its +// capacity is sufficient. This avoids a heap allocation for the regions slice +// on every AVC frame when called with a pre-allocated GfxHandler field. +func fillAVC420Stream(data []byte, out *avc420Stream) error { + if len(data) < 4 { + return fmt.Errorf("avc420 stream too short (%d bytes)", len(data)) + } + + numRegions := binary.LittleEndian.Uint32(data[:4]) + if numRegions > 65536 { + return fmt.Errorf("avc420: too many regions: %d", numRegions) + } + + // 4 bytes header + 10 bytes per region (8-byte rect + 2-byte quant/quality) + metaSize := 4 + int(numRegions)*10 + if metaSize > len(data) { + return fmt.Errorf("avc420: metadata truncated (need %d, have %d)", metaSize, len(data)) + } + + if cap(out.regions) >= int(numRegions) { + out.regions = out.regions[:numRegions] + } else { + out.regions = make([]avcRect, numRegions) + } + off := 4 + for i := range numRegions { + out.regions[i] = avcRect{ + left: binary.LittleEndian.Uint16(data[off:]), + top: binary.LittleEndian.Uint16(data[off+2:]), + right: binary.LittleEndian.Uint16(data[off+4:]), + bottom: binary.LittleEndian.Uint16(data[off+6:]), + } + off += 8 + } + out.h264Data = data[metaSize:] + return nil +} + +// parseAVC420Stream parses RDPGFX_AVC420_BITMAP_STREAM into a new struct. +// Callers that run on the decode goroutine should prefer fillAVC420Stream with +// a pre-allocated GfxHandler field to avoid per-frame heap allocations. +func parseAVC420Stream(data []byte) (*avc420Stream, error) { + var out avc420Stream + if err := fillAVC420Stream(data, &out); err != nil { + return nil, err + } + return &out, nil +} + +// parseAVC444Stream parses RDPGFX_AVC444_BITMAP_STREAM. +// Returns the main AVC420 stream, the auxiliary AVC420 stream, and the LC +// (luma-chroma) field. +// +// LC=0: both streams present; stream1 = main (YUV420), stream2 = chroma upgrade. +// LC=1: main stream only; stream2 is nil. +// LC=2: auxiliary only (chroma upgrade); stream1 is nil. +func parseAVC444Stream(data []byte) (stream1, stream2 *avc420Stream, lc uint8, err error) { + if len(data) < 4 { + return nil, nil, 0, fmt.Errorf("avc444 stream too short") + } + + cbField := binary.LittleEndian.Uint32(data[:4]) + lc = uint8((cbField >> 30) & 0x03) + cbStream1 := int(cbField & 0x3FFFFFFF) + rest := data[4:] + + switch lc { + case 0: // Both streams present + if cbStream1 > len(rest) { + return nil, nil, lc, fmt.Errorf("avc444: stream1 size %d exceeds data %d", cbStream1, len(rest)) + } + stream1, err = parseAVC420Stream(rest[:cbStream1]) + if err != nil { + return nil, nil, lc, err + } + if cbStream1 < len(rest) { + stream2, err = parseAVC420Stream(rest[cbStream1:]) + if err != nil { + slog.Debug("RDPGFX: AVC444 stream2 parse error (LC=0)", "err", err) + stream2 = nil + err = nil + } + } + return stream1, stream2, lc, nil + case 1: // Main stream only + streamData := rest + if cbStream1 > 0 && cbStream1 <= len(rest) { + streamData = rest[:cbStream1] + } + stream1, err = parseAVC420Stream(streamData) + return stream1, nil, lc, err + case 2: // Auxiliary only (chroma upgrade) + streamData := rest + if cbStream1 > 0 && cbStream1 <= len(rest) { + streamData = rest[:cbStream1] + } + stream2, err = parseAVC420Stream(streamData) + return nil, stream2, lc, err + default: + return nil, nil, lc, fmt.Errorf("avc444: invalid LC=%d", lc) + } +} + +// fillAVC444Stream parses data into g.avcStream1 and g.avcStream2, reusing +// their regions slices to avoid per-frame heap allocations. +// Safe: always called on the single decode goroutine. +func (g *GfxHandler) fillAVC444Stream(data []byte) (stream1, stream2 *avc420Stream, lc uint8, err error) { + if len(data) < 4 { + return nil, nil, 0, fmt.Errorf("avc444 stream too short") + } + + cbField := binary.LittleEndian.Uint32(data[:4]) + lc = uint8((cbField >> 30) & 0x03) + cbStream1 := int(cbField & 0x3FFFFFFF) + rest := data[4:] + + switch lc { + case 0: // Both streams present + if cbStream1 > len(rest) { + return nil, nil, lc, fmt.Errorf("avc444: stream1 size %d exceeds data %d", cbStream1, len(rest)) + } + if err = fillAVC420Stream(rest[:cbStream1], &g.avcStream1); err != nil { + return nil, nil, lc, err + } + stream1 = &g.avcStream1 + if cbStream1 < len(rest) { + if err2 := fillAVC420Stream(rest[cbStream1:], &g.avcStream2); err2 != nil { + slog.Debug("RDPGFX: AVC444 stream2 parse error (LC=0)", "err", err2) + } else { + stream2 = &g.avcStream2 + } + } + return stream1, stream2, lc, nil + case 1: // Main stream only + streamData := rest + if cbStream1 > 0 && cbStream1 <= len(rest) { + streamData = rest[:cbStream1] + } + if err = fillAVC420Stream(streamData, &g.avcStream1); err != nil { + return nil, nil, lc, err + } + return &g.avcStream1, nil, lc, nil + case 2: // Auxiliary only (chroma upgrade) + streamData := rest + if cbStream1 > 0 && cbStream1 <= len(rest) { + streamData = rest[:cbStream1] + } + if err = fillAVC420Stream(streamData, &g.avcStream2); err != nil { + return nil, nil, lc, err + } + return nil, &g.avcStream2, lc, nil + default: + return nil, nil, lc, fmt.Errorf("avc444: invalid LC=%d", lc) + } +} + +// avc444YPlane caches the tightly-packed luma plane (stride = Width) from the +// most recently decoded AVC444 main stream. It is used to combine with the +// auxiliary chroma stream when LC=2 frames arrive. +type avc444YPlane struct { + data []byte // luma Y, tight-packed, stride = w + u []byte // Cb (U) plane from stream1, half-res, stride = (w+1)/2 + v []byte // Cr (V) plane from stream1, half-res, stride = (w+1)/2 + stride int // = w + uvStride int // = (w+1)/2 + w, h int + fullRange bool + updatedAt time.Time // last time the cache was refreshed from a live main-stream decode +} + +// avc444YStaleness is the maximum age of the Y-plane cache before LC=2 +// combines are suppressed. When the main decoder (h264dec) stalls, the +// Y-plane is frozen while incoming LC=2 frames carry fresh chroma — combining +// stale luma with fresh chroma produces wrong colours. 500 ms is well above +// the inter-frame interval at typical RDP frame rates (≥2 fps) yet much lower +// than the 7-second hard stall threshold, so normal operation is unaffected. +const avc444YStaleness = 500 * time.Millisecond + +// avcHWStallQueueDepthHint is the queueDepth value reported in +// FRAME_ACKNOWLEDGE PDUs while the HW decoder is stalling (Y cache stale). +// Reporting a depth of 10 signals to the Windows RDP server that the client's +// decode backlog is growing, prompting it to reduce encoding quality and +// bitrate. This reduces the stream of LC=2 frames that accumulate during a +// VideoToolbox null-frame period and gives VT more headroom to flush its +// pipeline. The hint is cleared when the Y cache is refreshed (stall over). +const avcHWStallQueueDepthHint uint32 = 10 + +// isH264Keyframe returns true when data contains an IDR NAL unit (type 5), +// which marks the start of a new GOP (key frame). The scan handles both +// 3-byte (00 00 01) and 4-byte (00 00 00 01) Annex-B start codes. +func isH264Keyframe(data []byte) bool { + for i := 0; i+4 <= len(data); i++ { + // Look for Annex-B start code: 00 00 01 or 00 00 00 01. + if data[i] == 0x00 && data[i+1] == 0x00 { + var nalByte byte + if data[i+2] == 0x01 && i+3 < len(data) { + nalByte = data[i+3] + i += 2 + } else if data[i+2] == 0x00 && i+3 < len(data) && data[i+3] == 0x01 && i+4 < len(data) { + nalByte = data[i+4] + i += 3 + } else { + continue + } + nalType := nalByte & 0x1F + if nalType == 5 { // IDR slice + return true + } + } + } + return false +} + +// firstNALType returns the NAL unit type byte of the first Annex-B NAL in +// data, or 0xFF if none found. Useful for diagnosing decoder "buffering". +func firstNALType(data []byte) byte { + for i := 0; i+4 <= len(data); i++ { + if data[i] == 0x00 && data[i+1] == 0x00 { + if data[i+2] == 0x01 && i+3 < len(data) { + return data[i+3] & 0x1F + } else if data[i+2] == 0x00 && i+3 < len(data) && data[i+3] == 0x01 && i+4 < len(data) { + return data[i+4] & 0x1F + } + } + } + return 0xFF +} + + +// decoded frame plus the dirty rectangle list reported in the AVC420 stream +// header (in decoded-frame coordinates). When regions is non-empty callers +// can blit only those regions instead of the whole frame, which dramatically +// reduces per-frame copying for typical desktop video where most of the +// frame is unchanged from the previous frame. +// The pooled return value is true when the returned slice was acquired from +// bitmapBufPool; the caller must then call releaseBitmapBuf on it. +func (g *GfxHandler) decodeAVC420(data []byte, destX, destY, destW, destH int) ([]byte, []avcRect, bool) { + // Parse the stream header once and reuse for both the raw-NAL callback and + // the actual decode path, avoiding a redundant walk of the metadata. + parseErr := fillAVC420Stream(data, &g.avcStream1) + stream := &g.avcStream1 + // Compute isKF once — shared between the onH264Raw forwarding path and the + // maybeCacheStream1IDR call below, avoiding two linear scans of the NAL data. + var isKF bool + if parseErr == nil && len(stream.h264Data) > 0 { + isKF = isH264Keyframe(stream.h264Data) + } + if g.onH264Raw != nil && parseErr == nil && len(stream.h264Data) > 0 { + nalData := make([]byte, len(stream.h264Data)) + copy(nalData, stream.h264Data) + // AVC420(WTS1)语义:H264 流按整桌面尺寸编码,帧内"有效像素" + // 由元数据 regions 子矩形逐个刻画(实测子矩形之外散布色度清零 + // 的绿像素,连 PDU 大矩形内部也不例外——大矩形只是外框)。 + // 元数据缺失时回退 PDU 矩形。帧原点 (0,0)。 + var regions []int32 + if len(stream.regions) > 0 { + for _, r := range stream.regions { + regions = append(regions, int32(r.left), int32(r.top), int32(r.right), int32(r.bottom)) + } + } else { + regions = []int32{int32(destX), int32(destY), int32(destX + destW), int32(destY + destH)} + } + g.onH264Raw(0, 0, destX+destW, destY+destH, isKF, nalData, regions) + } + if g.h264dec == nil { + return nil, nil, false + } + if parseErr != nil { + slog.Warn("RDPGFX: AVC420 parse error", "err", parseErr) + return nil, nil, false + } + if len(stream.h264Data) == 0 { + return nil, nil, false + } + if isKF { + g.maybeCacheStream1IDR(stream.h264Data) + } + // For frames where only a small dirty area changed, pass region hints so + // the decoder can skip converting pixels outside those rectangles. This + // is safe here because decodeAVC420 uses blitAndEmitAVCRegions (which only + // reads dirty pixels) when shouldUseAVCRegions returns true. + if rh, ok := g.h264dec.(RegionHinter); ok && + len(stream.regions) > 0 && shouldUseAVCRegions(stream.regions, destW, destH) { + if cap(g.regionHintBuf) >= len(stream.regions) { + g.regionHintBuf = g.regionHintBuf[:len(stream.regions)] + } else { + g.regionHintBuf = make([][4]uint16, len(stream.regions)) + } + for i, r := range stream.regions { + g.regionHintBuf[i] = [4]uint16{r.left, r.top, r.right, r.bottom} + } + rh.SetRegionHint(g.regionHintBuf) + } + frame, err := g.h264dec.Decode(stream.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error", "err", err) + return nil, nil, false + } + if frame == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + slog.Debug("RDPGFX: H.264 decode returned nil frame (buffering?)") + return nil, nil, false + } + if frame.Dropped { + slog.Debug("RDPGFX: AVC420 frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + return nil, nil, false + } + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC420 decoded", "frameW", frame.Width, "frameH", frame.Height, "destW", destW, "destH", destH, "regions", len(stream.regions), "h264Len", len(stream.h264Data)) + } + g.noteSuccessfulDecode() + decoded, pooled := cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + return decoded, stream.regions, pooled +} + +// decodeAVC444 decodes AVC444 bitmap data to BGRA pixels. +// LC=0 and LC=1 decode the main YUV420 stream and cache the luma plane for +// potential LC=2 combine. LC=2 combines the cached luma with the auxiliary +// chroma stream decoded by the secondary decoder. +// The pooled return value is true when the returned slice was acquired from +// bitmapBufPool; the caller must then call releaseBitmapBuf on it. +func (g *GfxHandler) decodeAVC444(data []byte, destX, destY, destW, destH int) ([]byte, []avcRect, bool) { + // Parse the stream header once and reuse for both the raw-NAL callback and + // the actual decode path, avoiding a redundant walk of the metadata. + stream1, stream2, lc, parseErr := g.fillAVC444Stream(data) + // Compute isKF once — shared between the onH264Raw forwarding path and the + // maybeCacheStream1IDR call below, avoiding two linear scans of the NAL data. + var isKF bool + if parseErr == nil && stream1 != nil && len(stream1.h264Data) > 0 { + isKF = isH264Keyframe(stream1.h264Data) + } + if g.onH264Raw != nil && parseErr == nil && stream1 != nil && len(stream1.h264Data) > 0 { + nalData := make([]byte, len(stream1.h264Data)) + copy(nalData, stream1.h264Data) + // 同 AVC420:元数据子矩形优先(权威有效范围),缺失回退整表面矩形 + var regions []int32 + if len(stream1.regions) > 0 { + for _, r := range stream1.regions { + regions = append(regions, int32(r.left), int32(r.top), int32(r.right), int32(r.bottom)) + } + } else { + regions = []int32{int32(destX), int32(destY), int32(destX + destW), int32(destY + destH)} + } + g.onH264Raw(0, 0, destW, destH, isKF, nalData, regions) + } + if g.h264dec == nil { + return nil, nil, false + } + if parseErr != nil { + slog.Warn("RDPGFX: AVC444 parse error", "err", parseErr) + return nil, nil, false + } + if lc == 2 { + return g.decodeAVC444LC2(stream2, destW, destH) + } + if stream1 == nil || len(stream1.h264Data) == 0 { + return nil, nil, false + } + + // Pass region hints so the decoder skips converting pixels outside the + // dirty rectangles. Safe here because decodeAVC444 also uses + // blitAndEmitAVCRegions when shouldUseAVCRegions returns true. + if rh, ok := g.h264dec.(RegionHinter); ok && + len(stream1.regions) > 0 && shouldUseAVCRegions(stream1.regions, destW, destH) { + if cap(g.regionHintBuf) >= len(stream1.regions) { + g.regionHintBuf = g.regionHintBuf[:len(stream1.regions)] + } else { + g.regionHintBuf = make([][4]uint16, len(stream1.regions)) + } + for i, r := range stream1.regions { + g.regionHintBuf[i] = [4]uint16{r.left, r.top, r.right, r.bottom} + } + rh.SetRegionHint(g.regionHintBuf) + } + + var frame *H264Frame + var i420out *H264FrameI420 + var err error + isKeyFrame := isKF + if isKeyFrame { + // Cache IDR NAL data so the SW fallback decoder can be primed + // immediately after a VideoToolbox stall without waiting for VBox. + g.maybeCacheStream1IDR(stream1.h264Data) + } + isIDR := g.h264dec2 != nil && isKeyFrame + if isIDR { + // Reset per-GOP diagnostic flags so the LC=0 IDR and LC=2 combine + // after this IDR are sampled again for colour diagnostics. + g.lc2SampleLogged = false + g.lc2PFrameSampleLogged = false + g.lc0SampleLogged = false + } + if g.h264dec2 != nil { + // Cache luma for future LC=2 combine. + if i420dec, ok := g.h264dec.(I420Decoder); ok { + frame, i420out, err = i420dec.DecodeWithI420(stream1.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444)", "err", err) + return nil, nil, false + } + if i420out != nil { + g.updateAVC444YCache(i420out) + if isIDR { + // Snapshot the IDR luma separately. When a standalone + // LC=2 packet carries a stream2 IDR, the chroma data + // belongs to this GOP's first frame, so we must combine + // it with the IDR luma — not with a later P-frame's luma + // that has since overwritten avc444YPlane. + g.copyAVC444YToIDRCache() + } + } + } else { + frame, err = g.h264dec.Decode(stream1.h264Data) + } + } else { + frame, err = g.h264dec.Decode(stream1.h264Data) + } + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444)", "err", err) + return nil, nil, false + } + // Prime the aux decoder before any nil/drop checks so the stream2 IDR is + // never lost. On macOS, VideoToolbox returns nil frames for 1–3 s during + // initial warm-up; without this early call the stream2 IDR carried by the + // first LC=0 packet would be discarded (h264dec2 never created) and the + // renegotiation timer would degrade LC=2 to LC=0-only after retrying. + if lc == 0 && stream2 != nil && len(stream2.h264Data) > 0 { + g.primeAuxDecoder(stream2.h264Data) + } + if frame == nil { + if i420out == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + return nil, nil, false + } + // I420 fast path: HW decoder returned planar I420 instead of BGRA. + // Convert to BGRA using BT.709 (AVC444 standard encoding) so the + // BGRA rendering path can continue normally. + bgra, _ := i420ToBGRA(i420out) + if bgra == nil { + return nil, nil, false + } + frame = &H264Frame{Data: bgra, Width: i420out.Width, Height: i420out.Height} + } + if frame.Dropped { + slog.Debug("RDPGFX: AVC444 frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + // Touch Y cache timestamp so LC=2 can still combine with last valid luma + // while the server delivers a recovery IDR. + if g.avc444YPlane.w > 0 && !g.avc444YPlane.updatedAt.IsZero() { + g.avc444YPlane.updatedAt = time.Now() + } + return nil, nil, false + } + if !g.lc0SampleLogged && isIDR { + g.lc0SampleLogged = true + bgraData := frame.Data + w, h := frame.Width, frame.Height + for _, p := range [][2]int{{960, 400}, {480, 400}, {1440, 400}, {960, 600}, {100, 100}} { + px, py := p[0], p[1] + if px >= w || py >= h { + continue + } + off := (py*w + px) * 4 + if off+3 < len(bgraData) { + var rawY, rawU, rawV byte + if i420out != nil && py < i420out.Height && px < i420out.Width { + rawY = i420out.Y[py*i420out.YStride+px] + rawU = i420out.U[(py/2)*i420out.UStride+(px/2)] + rawV = i420out.V[(py/2)*i420out.VStride+(px/2)] + } + slog.Debug("H.264: pixel sample (LC=0 IDR frame)", + "x", px, "y", py, + "rawY", rawY, "rawU", rawU, "rawV", rawV, + "fullRange", i420out != nil && i420out.FullRange, + "B", bgraData[off], "G", bgraData[off+1], "R", bgraData[off+2]) + } + } + } + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC444 decoded", "frameW", frame.Width, "frameH", frame.Height, + "destW", destW, "destH", destH, "h264Len", len(stream1.h264Data)) + } + g.noteSuccessfulDecode() + decoded, pooled := cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + + return decoded, stream1.regions, pooled +} + +// decodeAVC420WithI420 decodes AVC420 bitmap data, returning BGRA pixels for +// the surface backing store and, when the underlying decoder supports I420 +// output, an optional H264FrameI420 for GPU-accelerated IYUV texture upload. +// i420 is nil when I420 extraction is unsupported or the frame dimensions are +// smaller than destW×destH. Callers must fall back to BGRA rendering when +// i420 is nil. +func (g *GfxHandler) decodeAVC420WithI420(data []byte, destX, destY, destW, destH int) (decoded []byte, i420 *H264FrameI420, regions []avcRect, pooled bool) { + if err := fillAVC420Stream(data, &g.avcStream1); err != nil { + slog.Warn("RDPGFX: AVC420 parse error", "err", err) + return + } + stream := &g.avcStream1 + if g.onH264Raw != nil && len(stream.h264Data) > 0 { + isKF := isH264Keyframe(stream.h264Data) + nalData := make([]byte, len(stream.h264Data)) + copy(nalData, stream.h264Data) + g.onH264Raw(destX, destY, destW, destH, isKF, nalData, nil) + } + if g.h264dec == nil || len(stream.h264Data) == 0 { + return + } + if isH264Keyframe(stream.h264Data) { + g.maybeCacheStream1IDR(stream.h264Data) + } + var frame *H264Frame + var err error + i420dec, hasI420 := g.h264dec.(I420Decoder) + if hasI420 { + var i420out *H264FrameI420 + frame, i420out, err = i420dec.DecodeWithI420(stream.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error", "err", err) + return + } + if i420out != nil && i420out.Width >= destW && i420out.Height >= destH { + i420 = i420out + } + } else { + frame, err = g.h264dec.Decode(stream.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error", "err", err) + return + } + } + // I420 fast path: frame is nil but i420 is non-nil — decoder produced output + // via the direct NV12/YUV420P copy path. Still counts as a successful decode. + if frame == nil && i420 == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + slog.Debug("RDPGFX: H.264 decode returned nil frame (buffering?)") + return + } + if frame != nil && frame.Dropped { + slog.Debug("RDPGFX: AVC420 (WithI420) frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + return + } + g.noteSuccessfulDecode() + if frame != nil { + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC420 decoded (WithI420)", "frameW", frame.Width, "frameH", frame.Height, + "destW", destW, "destH", destH, "hasI420", i420 != nil, + "regions", len(stream.regions), "h264Len", len(stream.h264Data)) + } + decoded, pooled = cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + } + regions = stream.regions + return +} + +// decodeAVC420WithNV12 decodes AVC420 bitmap data, returning native NV12 +// planes when the underlying decoder produces NV12 (typically VideoToolbox). +// If NV12 is unavailable, decoded may contain a BGRA fallback frame. +func (g *GfxHandler) decodeAVC420WithNV12(data []byte, destX, destY, destW, destH int) (decoded []byte, nv12 *H264FrameNV12, regions []avcRect, pooled bool) { + if err := fillAVC420Stream(data, &g.avcStream1); err != nil { + slog.Warn("RDPGFX: AVC420 parse error", "err", err) + return + } + stream := &g.avcStream1 + // Compute isKF once — shared between the onH264Raw forwarding path and the + // maybeCacheStream1IDR call below, avoiding two linear scans of the NAL data. + var isKF bool + if len(stream.h264Data) > 0 { + isKF = isH264Keyframe(stream.h264Data) + } + if g.onH264Raw != nil && len(stream.h264Data) > 0 { + nalData := make([]byte, len(stream.h264Data)) + copy(nalData, stream.h264Data) + g.onH264Raw(destX, destY, destW, destH, isKF, nalData, nil) + } + if g.h264dec == nil || len(stream.h264Data) == 0 { + return + } + if isKF { + g.maybeCacheStream1IDR(stream.h264Data) + } + var frame *H264Frame + var err error + nv12dec, hasNV12 := g.h264dec.(NV12Decoder) + if hasNV12 { + var nv12out *H264FrameNV12 + frame, nv12out, err = nv12dec.DecodeWithNV12(stream.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error", "err", err) + return + } + if nv12out != nil && nv12out.Width >= destW && nv12out.Height >= destH { + nv12 = nv12out + } + } else { + frame, err = g.h264dec.Decode(stream.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error", "err", err) + return + } + } + if frame == nil && nv12 == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + slog.Debug("RDPGFX: H.264 decode returned nil frame (buffering?)") + return + } + if frame != nil && frame.Dropped { + slog.Debug("RDPGFX: AVC420 (WithNV12) frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + return + } + g.noteSuccessfulDecode() + if frame != nil { + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC420 decoded (WithNV12)", "frameW", frame.Width, "frameH", frame.Height, + "destW", destW, "destH", destH, "hasNV12", nv12 != nil, + "regions", len(stream.regions), "h264Len", len(stream.h264Data)) + } + decoded, pooled = cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + } + regions = stream.regions + return +} + +// decodeAVC444WithI420 decodes AVC444 bitmap data, returning BGRA pixels and +// an optional I420 frame. LC=0 and LC=1 decode the main stream and cache the +// luma plane. LC=2 decodes the auxiliary chroma stream and combines it with +// the cached luma to produce BGRA; i420 is nil for LC=2 frames (GPU path falls +// back to BGRA). +func (g *GfxHandler) decodeAVC444WithI420(data []byte, destX, destY, destW, destH int) (decoded []byte, i420 *H264FrameI420, regions []avcRect, pooled bool) { + stream1, stream2, lc, err := g.fillAVC444Stream(data) + // Compute isKF once — shared between the onH264Raw forwarding path and the + // maybeCacheStream1IDR call below, avoiding two linear scans of the NAL data. + var isKF bool + if stream1 != nil && len(stream1.h264Data) > 0 { + isKF = isH264Keyframe(stream1.h264Data) + } + if g.onH264Raw != nil && stream1 != nil && len(stream1.h264Data) > 0 { + nalData := make([]byte, len(stream1.h264Data)) + copy(nalData, stream1.h264Data) + g.onH264Raw(destX, destY, destW, destH, isKF, nalData, nil) + } + if err != nil { + slog.Warn("RDPGFX: AVC444 parse error", "err", err) + return + } + if lc == 2 { + decoded, regions, pooled = g.decodeAVC444LC2(stream2, destW, destH) + return + } + if g.h264dec == nil || stream1 == nil || len(stream1.h264Data) == 0 { + return + } + if isKF { + g.maybeCacheStream1IDR(stream1.h264Data) + } + var frame *H264Frame + i420dec, hasI420 := g.h264dec.(I420Decoder) + if hasI420 { + var i420out *H264FrameI420 + frame, i420out, err = i420dec.DecodeWithI420(stream1.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444)", "err", err) + return + } + if i420out != nil { + if g.h264dec2 != nil { + g.updateAVC444YCache(i420out) + } + if i420out.Width >= destW && i420out.Height >= destH { + i420 = i420out + } + } + } else { + frame, err = g.h264dec.Decode(stream1.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444)", "err", err) + return + } + } + // Prime the aux decoder before checking frame.Dropped: stream2 IDR data + // must not be lost when the main frame is discarded due to zero-fill. + if lc == 0 && stream2 != nil && len(stream2.h264Data) > 0 { + g.primeAuxDecoder(stream2.h264Data) + } + // I420 fast path: frame is nil but i420 is non-nil — decoder produced output + // via the direct NV12/YUV420P copy path. Still counts as a successful decode. + if frame == nil && i420 == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + return + } + if frame != nil && frame.Dropped { + slog.Debug("RDPGFX: AVC444 (WithI420) frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + if g.avc444YPlane.w > 0 && !g.avc444YPlane.updatedAt.IsZero() { + g.avc444YPlane.updatedAt = time.Now() + } + return + } + g.noteSuccessfulDecode() + if frame != nil { + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC444 decoded (WithI420)", "frameW", frame.Width, "frameH", frame.Height, + "destW", destW, "destH", destH, "hasI420", i420 != nil, "h264Len", len(stream1.h264Data)) + } + decoded, pooled = cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + regions = stream1.regions + } + + return +} + +// decodeAVC444WithNV12 decodes AVC444 bitmap data, returning native NV12 planes +// for LC=0/LC=1 frames and BGRA for LC=2 chroma-combination frames. +// nv12 is non-nil only when the hardware decoder (VideoToolbox) produced NV12 +// output for a LC=0/LC=1 stream1 packet. LC=2 frames always return decoded +// BGRA with nv12==nil because the chroma supplement requires a CPU combine step. +func (g *GfxHandler) decodeAVC444WithNV12(data []byte, destX, destY, destW, destH int) (decoded []byte, nv12 *H264FrameNV12, regions []avcRect, pooled bool) { + stream1, stream2, lc, err := g.fillAVC444Stream(data) + var isKF bool + if stream1 != nil && len(stream1.h264Data) > 0 { + isKF = isH264Keyframe(stream1.h264Data) + } + if g.onH264Raw != nil && stream1 != nil && len(stream1.h264Data) > 0 { + nalData := make([]byte, len(stream1.h264Data)) + copy(nalData, stream1.h264Data) + g.onH264Raw(destX, destY, destW, destH, isKF, nalData, nil) + } + if err != nil { + slog.Warn("RDPGFX: AVC444 parse error", "err", err) + return + } + if lc == 2 { + decoded, regions, pooled = g.decodeAVC444LC2(stream2, destW, destH) + return + } + if g.h264dec == nil || stream1 == nil || len(stream1.h264Data) == 0 { + return + } + if isKF { + g.maybeCacheStream1IDR(stream1.h264Data) + } + isIDR := g.h264dec2 != nil && isKF + if isIDR { + g.lc2SampleLogged = false + g.lc2PFrameSampleLogged = false + g.lc0SampleLogged = false + } + var frame *H264Frame + nv12dec, hasNV12 := g.h264dec.(NV12Decoder) + if hasNV12 { + var nv12out *H264FrameNV12 + frame, nv12out, err = nv12dec.DecodeWithNV12(stream1.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444 WithNV12)", "err", err) + return + } + if nv12out != nil && nv12out.Width >= destW && nv12out.Height >= destH { + nv12 = nv12out + if g.h264dec2 != nil { + g.updateAVC444YCacheFromNV12(nv12out) + if isIDR { + g.copyAVC444YToIDRCache() + } + } + } + } else { + frame, err = g.h264dec.Decode(stream1.h264Data) + if err != nil { + slog.Warn("RDPGFX: H.264 decode error (AVC444 WithNV12 fallback)", "err", err) + return + } + } + // When the NV12 decoder returned a BGRA frame but no NV12 planes (e.g. the + // software decoder produced YUV420P instead of NV12), recover I420 from the + // decoder's side channel for Y cache update so LC=2 frames are not stalled. + if hasNV12 && nv12 == nil && frame != nil && g.h264dec2 != nil { + type i420LastDecoder interface { + LastI420() *H264FrameI420 + } + if p, ok := g.h264dec.(i420LastDecoder); ok { + if i420out := p.LastI420(); i420out != nil { + g.updateAVC444YCache(i420out) + if isIDR { + g.copyAVC444YToIDRCache() + } + } + } + } + // Prime aux decoder before nil/dropped checks so stream2 IDR data is never lost. + if lc == 0 && stream2 != nil && len(stream2.h264Data) > 0 { + g.primeAuxDecoder(stream2.h264Data) + } + if frame == nil && nv12 == nil { + g.maybeRequestKeyframe() + g.maybeNotifyDecoderBroken() + slog.Debug("RDPGFX: H.264 decode returned nil frame/nv12 (buffering?)") + return + } + if frame != nil && frame.Dropped { + slog.Debug("RDPGFX: AVC444 (WithNV12) frame intentionally dropped (zero-fill)") + g.trackSWFallbackDroppedFrame() + g.maybeRequestKeyframe() + if g.avc444YPlane.w > 0 && !g.avc444YPlane.updatedAt.IsZero() { + g.avc444YPlane.updatedAt = time.Now() + } + return + } + g.noteSuccessfulDecode() + if frame != nil { + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: AVC444 decoded (WithNV12)", "frameW", frame.Width, "frameH", frame.Height, + "destW", destW, "destH", destH, "hasNV12", nv12 != nil, "h264Len", len(stream1.h264Data)) + } + decoded, pooled = cropBGRA(frame.Data, frame.Width, frame.Height, destW, destH) + } + regions = stream1.regions + return +} + +// updateAVC444YCache copies the Y, U, and V planes from stream1's i420 into +// g.avc444YPlane for use when combining with an LC=2 auxiliary chroma frame. +// The U/V planes are stored half-res (stride = (w+1)/2) and provide the B2/B3 +// chroma values (even column, even row positions) that stream2 does not cover. +func (g *GfxHandler) updateAVC444YCache(i420 *H264FrameI420) { + w, h := i420.Width, i420.Height + uvStride := (w + 1) / 2 + uvH := (h + 1) / 2 + neededY := w * h + neededUV := uvStride * uvH + if cap(g.avc444YPlane.data) < neededY { + g.avc444YPlane.data = make([]byte, neededY) + } else { + g.avc444YPlane.data = g.avc444YPlane.data[:neededY] + } + if cap(g.avc444YPlane.u) < neededUV { + g.avc444YPlane.u = make([]byte, neededUV) + } else { + g.avc444YPlane.u = g.avc444YPlane.u[:neededUV] + } + if cap(g.avc444YPlane.v) < neededUV { + g.avc444YPlane.v = make([]byte, neededUV) + } else { + g.avc444YPlane.v = g.avc444YPlane.v[:neededUV] + } + // i420 planes are already tight-packed (strides == width/height from extractI420fromSrc). + // Run Y, U, V copies in parallel for large frames: each slice is an independent + // allocation so there is no aliasing between the goroutines' writes. + totalBytes := neededY + neededUV*2 + if totalBytes >= parallelConvertMinPixels*4 { + var wg sync.WaitGroup + wg.Add(3) + go func() { defer wg.Done(); copy(g.avc444YPlane.data, i420.Y) }() + go func() { defer wg.Done(); copy(g.avc444YPlane.u, i420.U) }() + go func() { defer wg.Done(); copy(g.avc444YPlane.v, i420.V) }() + wg.Wait() + } else { + copy(g.avc444YPlane.data, i420.Y) + copy(g.avc444YPlane.u, i420.U) + copy(g.avc444YPlane.v, i420.V) + } + g.avc444YPlane.stride = w + g.avc444YPlane.uvStride = uvStride + g.avc444YPlane.w = w + g.avc444YPlane.h = h + g.avc444YPlane.fullRange = i420.FullRange + g.avc444YPlane.updatedAt = time.Now() + // HW decoder is producing real frames again — clear any stall throttle so + // the server resumes its normal quality/bitrate. + g.SetQueueDepthHint(0) +} + +// updateAVC444YCacheFromNV12 updates the AVC444 Y/UV cache from a native NV12 +// frame produced by the hardware decoder (VideoToolbox on macOS). The NV12 +// interleaved UV plane is de-interleaved into separate U/V planes so that the +// cache layout matches what combineAVC444v2BGRA expects. +// +// fullRange is forced to false regardless of nv12.FullRange. VideoToolbox +// expands the H.264 limited-range chroma to full range when it outputs +// kCVPixelFormatType_420YpCbCr8BiPlanarFullRange, but the SDL2 Metal renderer +// applies limited-range BT.709 coefficients to the NV12 texture (SDL2 auto- +// selects BT.709 for HD resolutions). The LC=2 auxiliary-chroma stream is +// decoded by the SW decoder and keeps limited-range values. By always using +// fullRange=false here, combineAVC444v2BGRA applies the same limited-range +// BT.709 formula that SDL2 uses for the NV12 texture, eliminating the colour +// shift that was visible when the display alternated between LC=0 NV12 frames +// (SDL2 limited BT.709) and LC=2 BGRA overlay frames (previously BT.709 full- +// range). +func (g *GfxHandler) updateAVC444YCacheFromNV12(nv12 *H264FrameNV12) { + w, h := nv12.Width, nv12.Height + uvStride := (w + 1) / 2 + uvH := (h + 1) / 2 + neededY := w * h + neededUV := uvStride * uvH + if cap(g.avc444YPlane.data) < neededY { + g.avc444YPlane.data = make([]byte, neededY) + } else { + g.avc444YPlane.data = g.avc444YPlane.data[:neededY] + } + if cap(g.avc444YPlane.u) < neededUV { + g.avc444YPlane.u = make([]byte, neededUV) + } else { + g.avc444YPlane.u = g.avc444YPlane.u[:neededUV] + } + if cap(g.avc444YPlane.v) < neededUV { + g.avc444YPlane.v = make([]byte, neededUV) + } else { + g.avc444YPlane.v = g.avc444YPlane.v[:neededUV] + } + // Copy Y rows, respecting the source stride. + srcYStride := nv12.YStride + if srcYStride <= 0 { + srcYStride = w + } + for row := range h { + copy(g.avc444YPlane.data[row*w:row*w+w], nv12.Y[row*srcYStride:row*srcYStride+w]) + } + // De-interleave NV12 UV (interleaved UVUVUV…) into separate U/V planes. + srcUVStride := nv12.UVStride + if srcUVStride <= 0 { + srcUVStride = w + } + for row := range uvH { + srcRow := nv12.UV[row*srcUVStride : row*srcUVStride+w] + dstU := g.avc444YPlane.u[row*uvStride : row*uvStride+uvStride] + dstV := g.avc444YPlane.v[row*uvStride : row*uvStride+uvStride] + for col := range uvStride { + dstU[col] = srcRow[col*2] + dstV[col] = srcRow[col*2+1] + } + } + g.avc444YPlane.stride = w + g.avc444YPlane.uvStride = uvStride + g.avc444YPlane.w = w + g.avc444YPlane.h = h + // Force limited-range so combineAVC444v2BGRA uses the same BT.709 limited- + // range coefficients as the SDL2 Metal renderer (see comment above). + g.avc444YPlane.fullRange = false + g.avc444YPlane.updatedAt = time.Now() + g.SetQueueDepthHint(0) +} + +// copyAVC444YToIDRCache copies the current avc444YPlane content into +// avc444IDRYPlane. Called immediately after updating avc444YPlane from a +// stream1 IDR decode, so the IDR luma snapshot stays separate from any +// subsequent P-frame luma updates. +func (g *GfxHandler) copyAVC444YToIDRCache() { + src := &g.avc444YPlane + dst := &g.avc444IDRYPlane + if cap(dst.data) < len(src.data) { + dst.data = make([]byte, len(src.data)) + } else { + dst.data = dst.data[:len(src.data)] + } + if cap(dst.u) < len(src.u) { + dst.u = make([]byte, len(src.u)) + } else { + dst.u = dst.u[:len(src.u)] + } + if cap(dst.v) < len(src.v) { + dst.v = make([]byte, len(src.v)) + } else { + dst.v = dst.v[:len(src.v)] + } + // Parallel copy for large frames: each slice is a separate allocation. + totalBytes := len(src.data) + len(src.u) + len(src.v) + if totalBytes >= parallelConvertMinPixels*4 { + var wg sync.WaitGroup + wg.Add(3) + go func() { defer wg.Done(); copy(dst.data, src.data) }() + go func() { defer wg.Done(); copy(dst.u, src.u) }() + go func() { defer wg.Done(); copy(dst.v, src.v) }() + wg.Wait() + } else { + copy(dst.data, src.data) + copy(dst.u, src.u) + copy(dst.v, src.v) + } + dst.stride = src.stride + dst.uvStride = src.uvStride + dst.w = src.w + dst.h = src.h + dst.fullRange = src.fullRange + dst.updatedAt = src.updatedAt +} + +// maybeCacheStream1IDR stores h264Data as the latest stream1 IDR NAL data for +// later use when priming the SW fallback decoder after a VideoToolbox stall. +// The data is copied so the caller's buffer may be reused freely. +// Only call when isH264Keyframe(h264Data) is true. +func (g *GfxHandler) maybeCacheStream1IDR(h264Data []byte) { + g.lastStream1IDR = append(g.lastStream1IDR[:0], h264Data...) + g.lastStream1IDRTime = time.Now() + g.lastStream1IDRFrame = g.framesDecoded.Load() + if g.usingSWFallback { + // A natural IDR from the server arrived while in SW fallback mode. + // From this frame onwards the SW decoder has a fresh reference point and + // error concealment (block noise) should stop. This is the genuine + // resync point after a stale-IDR prime, so disarm the stale-prime + // corruption tracking: the primed frames were only suspect until a real + // IDR healed the picture. + g.swFallbackPrimed = false + g.swFallbackDroppedCount = 0 + slog.Debug("H.264: natural IDR received during SW fallback — block noise should stop", + "idrLen", len(h264Data), + "framesDecoded", g.lastStream1IDRFrame) + } else { + slog.Debug("H.264: stream1 IDR cached for SW fallback priming", + "idrLen", len(h264Data), + "framesDecoded", g.lastStream1IDRFrame) + } +} + +// isPlaneRegionBlank samples a 3×3 grid inside a rectangular region of a +// single plane and returns true when a majority of the samples are either +// near-zero (< loThreshold) or near-saturated (>= hiThreshold). It is used +// to detect the uninitialised/corrupt chroma states that produce green or +// pink overlays in AVC444v2 reconstruction. +func isPlaneRegionBlank(data []byte, stride, x0, y0, w, h int) bool { + if len(data) == 0 || w <= 0 || h <= 0 { + return false + } + const ( + loThreshold = 72 // below this: abnormally low (green monochrome) + hiThreshold = 235 // at or above this: near-saturation (pink overlay) + ) + nearZero, nearSat, total := 0, 0, 0 + for i := range 3 { + row := y0 + (i+1)*h/4 + if row < y0 || row >= y0+h { + continue + } + for j := range 3 { + col := x0 + (j+1)*w/4 + if col < x0 || col >= x0+w { + continue + } + total++ + v := data[row*stride+col] + if v < loThreshold { + nearZero++ + } else if v >= hiThreshold { + nearSat++ + } + } + } + if total == 0 { + return false + } + return nearZero*2 > total || nearSat*2 > total +} + +// isAuxChromaBlank returns true when any of the chroma-carrying planes in the +// stream2 auxiliary frame looks uninitialised or corrupt. In AVC444v2 the Y +// plane carries Cb (left half) and Cr (right half) for odd columns, while the +// U and V planes carry the remaining chroma positions for even columns on odd +// rows. Near-zero values in any of these planes produce a bright green frame; +// near-saturated values produce a pink/magenta overlay. Detecting this early +// lets decodeAVC444LC2 skip the combine and wait for real data. +// +// Two failure modes are detected: +// - Near-zero (< 20): codec not yet initialised; Windows Server initialises +// stream2 IDR with Cb≈0, Cr≈0 and sometimes emits sparse artefact pixels +// (Cb≈9–12) that a threshold of 8 would pass. Raising to 20 keeps those +// from triggering a combine that produces bright green blocks. +// - Near-saturation (≥ 235): indicates DPB mismatch or corruption in the aux +// decoder (h264dec2); a P-frame decoded against the wrong reference can +// produce near-maximal values, which encode Cb≈255/Cr≈255 and result in a +// pink/magenta overlay when combined with any luma. +func isAuxChromaBlank(f *H264FrameI420) bool { + if f == nil || f.Width < 16 || f.Height < 4 || len(f.Y) == 0 || len(f.U) == 0 || len(f.V) == 0 { + return false + } + w, h := f.Width, f.Height + halfW := w / 2 + uvW := halfW / 2 + uvH := (h + 1) / 2 + // Y plane left half: Cb for odd columns. + if isPlaneRegionBlank(f.Y, f.YStride, 0, 0, halfW, h) { + return true + } + // Y plane right half: Cr for odd columns. + if isPlaneRegionBlank(f.Y, f.YStride, halfW, 0, halfW, h) { + return true + } + // U plane: Cb/Cr for even columns on odd rows. + if isPlaneRegionBlank(f.U, f.UStride, 0, 0, uvW, uvH) { + return true + } + // V plane: Cb/Cr for even columns on odd rows. + if isPlaneRegionBlank(f.V, f.VStride, 0, 0, uvW, uvH) { + return true + } + return false +} + +// isAVC444YPlaneChromaBlank returns true when the cached stream1 chroma (U/V) +// looks corrupt. Green monochrome requires both Cb and Cr to collapse to +// near-zero, so the grid check requires both U and V at a sample point to be +// low. Near-saturation in both planes produces a pink/magenta overlay. +// +// The threshold (72) matches the low-chroma guard in the ffmpeg decoder plugin +// so a frame that poisoned the cache would also have been dropped there. +func isAVC444YPlaneChromaBlank(yp *avc444YPlane) bool { + if yp == nil || yp.w < 16 || yp.h < 4 || len(yp.u) == 0 || len(yp.v) == 0 { + return false + } + uvH := (yp.h + 1) / 2 + const ( + loThreshold = 72 // below this: near-zero (green monochrome) + hiThreshold = 235 // at or above this: near-saturation (pink overlay) + ) + nearZero, nearSat, total := 0, 0, 0 + for i := range 3 { + row := (i + 1) * uvH / 4 + if row >= uvH { + continue + } + for j := range 3 { + col := (j + 1) * yp.uvStride / 4 + if col >= yp.uvStride { + continue + } + total++ + u := yp.u[row*yp.uvStride+col] + v := yp.v[row*yp.uvStride+col] + if u < loThreshold && v < loThreshold { + nearZero++ + } else if u >= hiThreshold && v >= hiThreshold { + nearSat++ + } + } + } + if total == 0 { + return false + } + return nearZero*2 > total || nearSat*2 > total +} + +// combineAVC444v2BGRA implements the AVC444v2 chroma reconstruction defined in +// [MS-RDPEGFX 3.3.8.3.3] ("YUV420p Stream Combination for YUV444v2 mode"). +// +// Stream2 encodes the missing chroma positions that stream1's 4:2:0 quantiser +// discards, split across three "Bx areas" of the auxiliary I420 frame: +// +// B4/B5 — stream2 Y plane, each row: +// bytes [0, w/2) = Cb at all odd-x columns (U444[2k+1, y] for k=0..w/2-1) +// bytes [w/2, w) = Cr at all odd-x columns (V444[2k+1, y] for k=0..w/2-1) +// +// B6/B7 — stream2 U plane, each half-height row j: +// bytes [0, w/4) = Cb at even-x multiples of 4 (U444[4k, 2j+1]) +// bytes [w/4, w/2) = Cr at even-x multiples of 4 (V444[4k, 2j+1]) +// +// B8/B9 — stream2 V plane, each half-height row j: +// bytes [0, w/4) = Cb at even-x offset-2 cols (U444[4k+2, 2j+1]) +// bytes [w/4, w/2) = Cr at even-x offset-2 cols (V444[4k+2, 2j+1]) +// +// Positions not covered by stream2 (even-x, even-y) use stream1's half-res +// B2/B3 chroma values from the cached cachedU/cachedV planes. +// +// Parameters: +// +// yPlane/yStride – luma Y from stream1, tight-packed (stride=w) +// cachedU/cachedV – Cb/Cr from stream1, half-res (stride=uvStride=(w+1)/2) +// i420aux – I420 output from decoding stream2 +// fullRange – true for PC-range [0-255], false for video [16-235] +// maxConvertWorkers caps the number of goroutines used to parallelise the +// per-row YCbCr→BGRA conversions. Beyond ~8 workers the conversion is limited +// by memory bandwidth rather than CPU, so additional workers only add +// scheduling overhead without speeding up the conversion. +const maxConvertWorkers = 8 + +// parallelConvertMinPixels is the frame-area threshold below which conversion +// runs serially: for small frames the goroutine spawn/join overhead exceeds the +// work saved by splitting the rows across cores. +const parallelConvertMinPixels = 256 * 256 + +// parallelRows splits the row range [0,h) into up to maxConvertWorkers +// contiguous chunks and runs fn(y0,y1) for each chunk concurrently, returning +// only once every chunk has finished. Each chunk writes a disjoint set of +// output rows and reads the shared input planes read-only, so the chunks are +// data-race free and the combined result is identical to a serial run. +// +// For small frames (area < parallelConvertMinPixels) fn is invoked once over +// the full range on the calling goroutine, avoiding goroutine overhead. +func parallelRows(w, h int, fn func(y0, y1 int)) { + workers := runtime.GOMAXPROCS(0) + if workers > maxConvertWorkers { + workers = maxConvertWorkers + } + if workers <= 1 || h < 2 || w*h < parallelConvertMinPixels { + fn(0, h) + return + } + if workers > h { + workers = h + } + chunk := (h + workers - 1) / workers + var wg sync.WaitGroup + for y0 := 0; y0 < h; y0 += chunk { + y1 := y0 + chunk + if y1 > h { + y1 = h + } + wg.Add(1) + go func(a, b int) { + defer wg.Done() + fn(a, b) + }(y0, y1) + } + wg.Wait() +} + +// combineAVC444v2BGRA combines luma from stream1 with per-pixel chroma from +// stream1 and stream2 to produce a BGRA frame. The fullRange branch is hoisted +// outside the inner loop; row-level offsets are computed once per row. The row +// range is split across cores via parallelRows for large frames. +// +// When dirtyRegions is non-nil, only rows covered by at least one region are +// converted; all other rows in the output buffer are left as pool garbage. +// Callers must only pass non-nil dirtyRegions when shouldUseAVCRegions is true, +// because in that case blitAndEmitAVCRegions reads only within the dirty rects +// and never accesses the uninitialised rows. +func combineAVC444v2BGRA( + yPlane []byte, yStride int, + cachedU, cachedV []byte, uvStride int, + i420aux *H264FrameI420, + fullRange bool, + w, h int, + dirtyRegions []avcRect, +) (out []byte, pooled bool) { + if len(yPlane) == 0 || len(cachedU) == 0 || len(cachedV) == 0 || w <= 0 || h <= 0 { + return nil, false + } + if i420aux == nil || len(i420aux.Y) == 0 || len(i420aux.U) == 0 || len(i420aux.V) == 0 { + return nil, false + } + out = acquireBitmapBuf(w * h * 4) + halfW := w / 2 + quarterW := w / 4 + auxYStride := i420aux.YStride + auxUStride := i420aux.UStride + auxVStride := i420aux.VStride + + // Build per-row dirty mask when only partial conversion is needed. Rows not + // covered by any dirty region are skipped inside rowFn; the corresponding + // output bytes remain as pool-buffer data that blitAndEmitAVCRegions never + // reads (it only accesses pixels within the dirty rectangles). + var rowDirty []bool + if len(dirtyRegions) > 0 { + rowDirty = make([]bool, h) + for _, r := range dirtyRegions { + r0 := max(0, int(r.top)) + r1 := min(h, int(r.bottom)) + for y := r0; y < r1; y++ { + rowDirty[y] = true + } + } + } + + // Split on fullRange once so the inner loop body is branch-free for the + // YCbCr→BGRA conversion coefficients. Each output row is independent, so + // parallelRows splits the rows across cores for large frames. + var rowFn func(y0, y1 int) + if fullRange { + rowFn = func(y0, y1 int) { + for row := y0; row < y1; row++ { + if rowDirty != nil && !rowDirty[row] { + continue + } + yRowOff := row * yStride + uvRow := row >> 1 + uvRowOff := uvRow * uvStride + auxYRowOff := row * auxYStride + auxURowOff := uvRow * auxUStride + auxVRowOff := uvRow * auxVStride + outIdx := row * w * 4 + for col := range w { + Y := yPlane[yRowOff+col] + var Cb, Cr byte + if col&1 == 1 { + k := col >> 1 + Cb = i420aux.Y[auxYRowOff+k] + Cr = i420aux.Y[auxYRowOff+halfW+k] + } else if row&1 == 0 { + k := col >> 1 + Cb = cachedU[uvRowOff+k] + Cr = cachedV[uvRowOff+k] + } else { + k := col >> 2 + if col&2 == 0 { + Cb = i420aux.U[auxURowOff+k] + Cr = i420aux.U[auxURowOff+quarterW+k] + } else { + Cb = i420aux.V[auxVRowOff+k] + Cr = i420aux.V[auxVRowOff+quarterW+k] + } + } + y := int(Y) + u := int(Cb) - 128 + v := int(Cr) - 128 + out[outIdx] = clampByte((256*y + 475*u + 128) >> 8) + out[outIdx+1] = clampByte((256*y - 48*u - 120*v + 128) >> 8) + out[outIdx+2] = clampByte((256*y + 403*v + 128) >> 8) + out[outIdx+3] = 255 + outIdx += 4 + } + } + } + } else { + rowFn = func(y0, y1 int) { + for row := y0; row < y1; row++ { + if rowDirty != nil && !rowDirty[row] { + continue + } + yRowOff := row * yStride + uvRow := row >> 1 + uvRowOff := uvRow * uvStride + auxYRowOff := row * auxYStride + auxURowOff := uvRow * auxUStride + auxVRowOff := uvRow * auxVStride + outIdx := row * w * 4 + for col := range w { + Y := yPlane[yRowOff+col] + var Cb, Cr byte + if col&1 == 1 { + k := col >> 1 + Cb = i420aux.Y[auxYRowOff+k] + Cr = i420aux.Y[auxYRowOff+halfW+k] + } else if row&1 == 0 { + k := col >> 1 + Cb = cachedU[uvRowOff+k] + Cr = cachedV[uvRowOff+k] + } else { + k := col >> 2 + if col&2 == 0 { + Cb = i420aux.U[auxURowOff+k] + Cr = i420aux.U[auxURowOff+quarterW+k] + } else { + Cb = i420aux.V[auxVRowOff+k] + Cr = i420aux.V[auxVRowOff+quarterW+k] + } + } + c := int(Y) - 16 + u := int(Cb) - 128 + v := int(Cr) - 128 + out[outIdx] = clampByte((298*c + 541*u + 128) >> 8) + out[outIdx+1] = clampByte((298*c - 55*u - 136*v + 128) >> 8) + out[outIdx+2] = clampByte((298*c + 459*v + 128) >> 8) + out[outIdx+3] = 255 + outIdx += 4 + } + } + } + } + parallelRows(w, h, rowFn) + return out, true +} + +// i420ToBGRA converts a planar I420 frame to a packed BGRA buffer using BT.709 +// coefficients (matching AVC444 content encoding). Used when the I420 fast path +// is active and a BGRA output is required by the rendering path. +// +// Optimised: the fullRange branch is hoisted outside both loops so the inner +// loop body is branch-free, row offsets are computed once per row, and outIdx +// advances by 4 instead of recomputing col*4 per pixel. +func i420ToBGRA(src *H264FrameI420) ([]byte, bool) { + if src == nil || src.Width <= 0 || src.Height <= 0 { + return nil, false + } + w, h := src.Width, src.Height + out := acquireBitmapBuf(w * h * 4) + var rowFn func(y0, y1 int) + if src.FullRange { + rowFn = func(yStart, yEnd int) { + for row := yStart; row < yEnd; row++ { + yOff := row * src.YStride + uvOff := (row >> 1) * src.UStride + uvOffV := (row >> 1) * src.VStride + outIdx := row * w * 4 + for col := range w { + y := int(src.Y[yOff+col]) + uv := col >> 1 + u := int(src.U[uvOff+uv]) - 128 + v := int(src.V[uvOffV+uv]) - 128 + out[outIdx] = clampByte((256*y + 475*u + 128) >> 8) + out[outIdx+1] = clampByte((256*y - 48*u - 120*v + 128) >> 8) + out[outIdx+2] = clampByte((256*y + 403*v + 128) >> 8) + out[outIdx+3] = 255 + outIdx += 4 + } + } + } + } else { + rowFn = func(yStart, yEnd int) { + for row := yStart; row < yEnd; row++ { + yOff := row * src.YStride + uvOffU := (row >> 1) * src.UStride + uvOffV := (row >> 1) * src.VStride + outIdx := row * w * 4 + for col := range w { + c := int(src.Y[yOff+col]) - 16 + uv := col >> 1 + u := int(src.U[uvOffU+uv]) - 128 + v := int(src.V[uvOffV+uv]) - 128 + out[outIdx] = clampByte((298*c + 541*u + 128) >> 8) + out[outIdx+1] = clampByte((298*c - 55*u - 136*v + 128) >> 8) + out[outIdx+2] = clampByte((298*c + 459*v + 128) >> 8) + out[outIdx+3] = 255 + outIdx += 4 + } + } + } + } + parallelRows(w, h, rowFn) + return out, true +} + +// avc444bt709BGRA converts one YCbCr pixel to BGRA using BT.709 coefficients, +// matching FreeRDP's general_YUV444ToBGRX implementation. +// Windows AVC444v2 content is encoded in BT.709; using BT.601 here was the +// cause of red color bleeding on LC=2 chroma-upgrade frames. +// Cb and Cr are raw (0-255); the function subtracts 128 internally. +// +// Full range (Y∈[0,255]): R = Y + 1.5748*(Cr-128) ≈ (256y + 403v) >> 8 +// Limited range (Y∈[16,235]): R = 1.164*(Y-16) + 1.793*(Cr-128) ≈ (298c + 459v) >> 8 +func avc444bt709BGRA(Y, Cb, Cr byte, fullRange bool, dst []byte) { + u := int(Cb) - 128 + v := int(Cr) - 128 + var r, g, b int + if fullRange { + y := int(Y) + r = (256*y + 403*v + 128) >> 8 + g = (256*y - 48*u - 120*v + 128) >> 8 + b = (256*y + 475*u + 128) >> 8 + } else { + c := int(Y) - 16 + r = (298*c + 459*v + 128) >> 8 + g = (298*c - 55*u - 136*v + 128) >> 8 + b = (298*c + 541*u + 128) >> 8 + } + dst[0] = clampByte(b) + dst[1] = clampByte(g) + dst[2] = clampByte(r) + dst[3] = 255 +} + +// clampByte clamps an integer to [0, 255] using branchless min/max built-ins. +func clampByte(v int) byte { + return byte(max(0, min(255, v))) +} + +// primeAuxDecoder feeds stream2 data from an LC=0 packet to h264dec2 so that +// the decoder's decoded-picture buffer (DPB) stays in sync with the full +// stream2 H.264 sequence. Stream2 frames are always part of one continuous +// H.264 sequence: the IDR is carried in LC=0 (and duplicated in a standalone +// LC=2 packet), and subsequent P-frames arrive via BOTH LC=0 packets and +// standalone LC=2 packets. If primeAuxDecoder only decoded IDRs, h264dec2's +// DPB would be stuck at the IDR while the server advanced the sequence through +// several LC=0 P-frames; the first standalone LC=2 P-frame would then be +// decoded against the wrong reference, producing all-zero chroma (Cb=0, +// Cr=0) and a full-screen green tint. By decoding ALL stream2 frames here +// (output discarded), h264dec2's DPB is always at the correct reference when +// primeH264dec2KeepDPB feeds stream2 data to h264dec2 and discards the +// output. Call this whenever decodeAVC444LC2 must skip the combine step +// (Y cache empty or stale) so that h264dec2's decoded-picture buffer stays +// in sync with the stream2 H.264 sequence. Without this, the next +// standalone LC=2 P-frame would reference a DPB state that is behind +// the expected position, causing FFmpeg to produce all-zero chroma +// (Cb=0, Cr=0) and a full-screen green tint. +func (g *GfxHandler) primeH264dec2KeepDPB(h264Data []byte) { + if g.h264dec2 == nil { + return + } + i420dec, ok := g.h264dec2.(I420Decoder) + if !ok { + return + } + _, _, err := i420dec.DecodeWithI420(h264Data) + if err != nil { + slog.Debug("RDPGFX: LC=2 DPB prime error", "err", err) + } + if g.h264dec2 != nil && g.h264dec2.IsBroken() { + g.h264dec2.Close() + g.h264dec2 = nil + g.startAuxDecoderBrokenTimer() + } +} + +// decodeAVC444LC2 decodes a standalone LC=2 P-frame. +func (g *GfxHandler) primeAuxDecoder(h264Data []byte) { + // Mark that stream2 data has appeared in an LC=0 packet. VirtualBox VRDE + // never includes stream2, so this flag distinguishes VirtualBox from Windows. + g.stream2EverSeen = true + isIDR := h264PacketHasIDR(h264Data) + if g.h264dec2 == nil { + if !isIDR { + // No aux decoder yet; wait for the stream2 IDR to create one. + return + } + // A stream2 IDR arrived — clear any permanent-degrade state so LC=2 + // can recover (e.g. after a server-side GOP reset much later in the session). + if g.lc2PermanentlyDegraded { + slog.Debug("H.264: stream2 IDR received after LC=2 degrade — recovering aux decoder") + g.lc2PermanentlyDegraded = false + g.auxDecoderNoIDRRetries = 0 + } + // Recreate aux decoder on a stream2 IDR so it starts with a clean + // reference frame. This avoids the rapid create/destroy cycle that + // can destabilise the decoder. + slog.Debug("H.264: recreating aux decoder on stream2 IDR") + g.h264dec2 = newH264DecoderSW() + g.stopAuxDecoderBrokenTimer() // LC=0 IDR arrived; cancel recovery timer + // Fall through to prime the freshly-created decoder with this IDR. + } + // If the aux decoder is broken, reset it only on an IDR (P-frames cannot + // start a new decode sequence). + if g.h264dec2.IsBroken() { + if !isIDR { + return + } + // Stream2 IDR received while aux decoder is broken — recreate it now + // and fall through to prime the fresh decoder with this IDR. + // (Previously this closed and waited for a *second* IDR which often + // never arrived, permanently losing LC=2 quality for the session.) + slog.Debug("H.264: recreating broken aux decoder on stream2 IDR") + g.h264dec2.Close() + g.h264dec2 = newH264DecoderSW() + g.stopAuxDecoderBrokenTimer() + } + i420dec, ok := g.h264dec2.(I420Decoder) + if !ok { + return + } + _, i420primed, err := i420dec.DecodeWithI420(h264Data) + if err != nil { + slog.Debug("RDPGFX: AVC444 aux prime error", "err", err) + } + // The pre-flight stall detector inside DecodeWithI420 can set broken=true + // and return nil,nil without an error (broken state invisible to caller). + // Check IsBroken() after the call to catch this case. + if g.h264dec2.IsBroken() { + slog.Debug("H.264: aux decoder broken after prime, waiting for IDR to recreate") + g.h264dec2.Close() + g.h264dec2 = nil + g.startAuxDecoderBrokenTimer() + return + } + // For P-frames, validate the decoded output. If the primed output looks + // blank (near-zero or near-saturated chroma), the DPB is likely corrupted + // (e.g. due to a dropped LC=0 PDU that left h264dec2 out of sync). + // Reset h264dec2 immediately so the DPB corruption does not cascade into + // the subsequent LC=2 standalone decode. The IDR case is excluded because + // near-zero output is expected during codec initialisation. + if !isIDR && i420primed != nil && isAuxChromaBlank(i420primed) { + slog.Debug("H.264: aux decoder DPB desynced during priming (P-frame blank chroma), resetting") + g.h264dec2.Close() + g.h264dec2 = nil + g.startAuxDecoderBrokenTimer() + } +} + +// decodeAVC444LC2 decodes an AVC444 LC=2 chroma-upgrade frame. +// It decodes stream2 via the auxiliary decoder, then combines the cached luma +// (Y plane) with the auxiliary chroma (Y2 = U/Cb channel, U2 = V/Cr channel) +// to produce a BGRA frame. +func (g *GfxHandler) decodeAVC444LC2(stream2 *avc420Stream, destW, destH int) (decoded []byte, regions []avcRect, pooled bool) { + // Record LC=2 arrival unconditionally so maybeRenegotiateCapabilities can + // distinguish an active-LC=2-only server from a truly idle server. + g.lastLC2RecvTime.Store(time.Now().UnixNano()) + if g.h264dec2 == nil { + if g.lc2PermanentlyDegraded { + // Server has proven it won't deliver stream2 IDRs; skip silently + // without arming the timer to avoid an endless renegotiation loop. + return + } + // If this standalone LC=2 frame carries an IDR, use it to create and + // prime h264dec2 directly. Some servers deliver the ForceRefresh IDR + // response as LC=1 (luma only) rather than LC=0 (both streams), so the + // IDR in the "duplicate" standalone LC=2 packet is the only opportunity + // to initialise the aux decoder without a full reconnect. + if stream2 != nil && len(stream2.h264Data) > 0 && isH264Keyframe(stream2.h264Data) { + slog.Debug("H.264: creating aux decoder from standalone LC=2 IDR") + g.h264dec2 = newH264DecoderSW() + g.stopAuxDecoderBrokenTimer() + g.auxDecoderNoIDRRetries = 0 + // Fall through to the decode path below. + } else { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (no aux decoder)") + // Arm the renegotiation timer so maybeRenegotiateCapabilities fires if + // no stream2 IDR arrives to prime h264dec2 within auxDecoderBrokenTimeout. + // This is idempotent — subsequent calls are no-ops while the timer runs. + g.startAuxDecoderBrokenTimer() + return + } + } + if stream2 == nil || len(stream2.h264Data) == 0 { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (empty aux stream)") + return + } + // If the main decoder is broken (e.g. HW stall or no IDR received), trigger + // soft reset so it can recover even when only LC=2 (chroma-only) frames are + // arriving and the LC=0/1 decode path never gets called. + if g.h264dec != nil && g.h264dec.IsBroken() { + g.maybeNotifyDecoderBroken() + return + } + if g.avc444YPlane.w == 0 { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (no cached luma)") + // Still advance h264dec2's DPB so the next standalone LC=2 P-frame + // finds the correct reference. Without this the DPB falls behind and + // FFmpeg outputs all-zero chroma (green tiles) on the next LC=2 decode. + g.primeH264dec2KeepDPB(stream2.h264Data) + g.maybeRequestKeyframe() + return + } + // Skip the combine when the Y cache is stale: the main decoder is likely + // stalling (VideoToolbox null frames). Combining old luma with fresh chroma + // produces visible colour artefacts. We suppress LC=2 output until h264dec + // delivers a fresh frame and refreshes the cache. + if !g.avc444YPlane.updatedAt.IsZero() && time.Since(g.avc444YPlane.updatedAt) > avc444YStaleness { + age := time.Since(g.avc444YPlane.updatedAt).Round(time.Millisecond) + slog.Debug("RDPGFX: AVC444 LC=2 skipped (Y cache stale, main decoder likely stalling)", + "age", age) + // Advance h264dec2's DPB even though we skip the combine, so that it + // stays in sync with the stream2 sequence and recovers cleanly once the + // main decoder exits its stall. + g.primeH264dec2KeepDPB(stream2.h264Data) + // Signal the server to reduce encoding quality/bitrate while the HW + // decoder is stalling. This throttles the stream of LC=2 frames that + // accumulate during VideoToolbox null-frame periods and gives VT more + // headroom to flush its pipeline. The hint is cleared in + // updateAVC444YCache when the HW decoder resumes real-frame output. + g.SetQueueDepthHint(avcHWStallQueueDepthHint) + // During a VideoToolbox stall h264dec.NeedsKeyframe() is false (the + // decoder has not been reset) so maybeRequestKeyframe() returns early. + // Request a keyframe directly here, reusing the shared rate-limiter, so + // the server delivers a fresh IDR that can help break the VT stall. + const keyframeRequestInterval = 2 * time.Second + if g.onKeyframeRequest != nil && time.Since(g.lastKeyframeRequest) >= keyframeRequestInterval { + g.lastKeyframeRequest = time.Now() + go g.onKeyframeRequest() + } + return + } + i420dec, ok := g.h264dec2.(I420Decoder) + if !ok { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (aux decoder lacks I420 support)") + return + } + _, i420aux, err := i420dec.DecodeWithI420(stream2.h264Data) + if err != nil { + slog.Warn("RDPGFX: AVC444 LC=2 aux decode error", "err", err) + if g.h264dec2.IsBroken() { + g.h264dec2.Close() + g.h264dec2 = nil + g.startAuxDecoderBrokenTimer() + } + return + } + if i420aux == nil { + slog.Debug("RDPGFX: AVC444 LC=2 aux decode buffering", + "h264Len", len(stream2.h264Data), + "firstNAL", firstNALType(stream2.h264Data), + "isIDR", isH264Keyframe(stream2.h264Data)) + // The pre-flight stall detector inside Decode() may have set broken=true + // and returned nil without an error. Detect and tear down here; the + // decoder will be recreated by primeAuxDecoder when the next stream2 + // IDR arrives, avoiding a rapid VT session create/destroy cycle. + if g.h264dec2 != nil && g.h264dec2.IsBroken() { + slog.Debug("H.264: aux decoder broken during LC=2 decode, waiting for IDR to recreate") + g.h264dec2.Close() + g.h264dec2 = nil + // Do NOT call maybeRequestKeyframe() here: ForceRefresh only delivers + // LC=1 luma IDR, not a stream2/chroma IDR. h264dec2 will be re-primed + // naturally when the next LC=0 frame arrives via primeAuxDecoder. + // The aux decoder broken timer will escalate to caps renegotiation if + // no LC=0 IDR arrives within auxDecoderBrokenTimeout. + g.startAuxDecoderBrokenTimer() + } + return + } + // Detect invalid aux chroma: two failure modes trigger this check. + // 1. Near-zero (Cb≈0, Cr≈0): Windows Server initialises stream2 IDR with + // Y≈0 and only refreshes regions that change; combining zero chroma with + // any luma produces BGRA(0,135,0,255) — a bright green screen. + // 2. Near-saturation (Cb≈255 or Cr≈255): DPB mismatch or aux decoder + // corruption that produces near-maximal stream2 Y values; these encode + // as extreme chroma and produce a pink/magenta overlay when combined. + // Determine IDR status before the blank-chroma check so it can drive the + // h264dec2 reset decision below. + stream2IsIDR := isH264Keyframe(stream2.h264Data) + if isAuxChromaBlank(i420aux) { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (stream2 chroma invalid: near-zero or near-saturated)") + // For P-frames, corrupt chroma means h264dec2's DPB has diverged from + // the server's reference (typically from a dropped LC=0 PDU). Decoding + // further P-frames against this wrong DPB would produce equally wrong + // output on every subsequent LC=2, perpetuating the pink/green artefact. + // Reset h264dec2 now so the DPB corruption does not cascade; recovery + // will happen automatically on the next stream2 IDR arriving in an LC=0. + // IDRs are excluded because near-zero chroma is expected at GOP start + // during stream2 codec initialisation and should not trigger a reset. + if !stream2IsIDR && g.h264dec2 != nil { + slog.Debug("H.264: aux decoder reset after P-frame blank chroma (DPB cascade prevention)") + g.h264dec2.Close() + g.h264dec2 = nil + g.startAuxDecoderBrokenTimer() + } + return + } + // Select the luma plane for the combine. When stream2 carries an IDR its + // chroma data corresponds to the GOP-boundary frame, not to the latest + // P-frame. Using avc444IDRYPlane (a snapshot of the luma at the moment + // stream1's IDR was decoded) avoids combining mismatched luma/chroma planes + // and eliminates the transient green tint that appears at GOP boundaries + // when the server delivers the stream2 IDR as a standalone LC=2 packet. + // Fall back to avc444YPlane when no IDR snapshot is available (e.g. the + // VideoToolbox pipeline delayed the IDR output past the P-frame boundary). + yp := &g.avc444YPlane + if stream2IsIDR && g.avc444IDRYPlane.w > 0 { + yp = &g.avc444IDRYPlane + slog.Debug("RDPGFX: AVC444 LC=2 IDR combine using IDR luma snapshot") + } + w, h := yp.w, yp.h + if i420aux.Width < w || i420aux.Height < h { + slog.Debug("RDPGFX: AVC444 LC=2 aux frame too small", + "auxW", i420aux.Width, "auxH", i420aux.Height, "lumaW", w, "lumaH", h) + return + } + // Guard against corrupt cached stream1 chroma. A frame that slipped past + // the decoder's low-chroma guard can poison the U/V cache; combining that + // with any stream2 chroma produces green/pink artefacts on the even-column, + // even-row pixels. Skip the combine and ask for a fresh IDR. + if isAVC444YPlaneChromaBlank(yp) { + slog.Debug("RDPGFX: AVC444 LC=2 skipped (cached stream1 chroma blank/corrupt)") + g.primeH264dec2KeepDPB(stream2.h264Data) + g.maybeRequestKeyframe() + return + } + // Pass dirty regions to combineAVC444v2BGRA so it can skip unchanged rows + // (significant savings for frames where only a small area updates). + // Only do this when shouldUseAVCRegions is true: in that case the caller + // (decodeAVC444 / decodeAVC444WithI420) routes the output through + // blitAndEmitAVCRegions, which reads only within the dirty rectangles, so + // any uninitialized rows in the output buffer are never accessed. + // + // Use destW/destH (surface dimensions) for the shouldUseAVCRegions check, + // not w/h (decoded frame dimensions). The region coordinates are in surface + // space, and the callers (WTS1/WTS2) also call shouldUseAVCRegions with + // surface dimensions. Using different dimensions here could cause + // combineRegions to be set (skipping rows, leaving stale pool garbage) while + // the caller falls through to blitToSurface (reading all rows) — writing + // that garbage to the display. Using destW/destH keeps the two decisions + // in sync and is more correct since the regions are in surface coordinate space. + var combineRegions []avcRect + if len(stream2.regions) > 0 && shouldUseAVCRegions(stream2.regions, destW, destH) { + combineRegions = stream2.regions + } + combined, _ := combineAVC444v2BGRA( + yp.data, yp.stride, + yp.u, yp.v, yp.uvStride, + i420aux, + yp.fullRange, + w, h, + combineRegions, + ) + if combined == nil { + return + } + // Mark that LC=2 has produced at least one frame this session. + // maybeRenegotiateCapabilities uses this to distinguish "was working then broke" + // (needs reconnect) from "never worked" (graceful LC=0 degradation). + g.lc2EverDecoded = true + g.auxDecoderNoIDRRetries = 0 // reset so a future break starts retries from scratch + // lc2Sample logs the actual Cb/Cr values used by combineAVC444v2BGRA for + // position (px,py), which depend on the B-area that pixel falls into. + halfW := w / 2 + quarterW := w / 4 + lc2Sample := func(px, py int) { + if px >= w || py >= h { + return + } + off := (py*w + px) * 4 + if off+3 >= len(combined) { + return + } + uvRow := py >> 1 + var actualCb, actualCr byte + var barea string + if px&1 == 1 { + // B4/B5: odd column — Cb/Cr packed in stream2 Y plane. + barea = "B4/B5" + k := px >> 1 + auxYRow := i420aux.Y[py*i420aux.YStride:] + actualCb = auxYRow[k] + actualCr = auxYRow[halfW+k] + } else if py&1 == 0 { + // B2/B3: even column, even row — from stream1 cached chroma. + barea = "B2/B3" + actualCb = yp.u[uvRow*yp.uvStride+(px>>1)] + actualCr = yp.v[uvRow*yp.uvStride+(px>>1)] + } else { + k2 := px >> 2 + if px&2 == 0 { + // B6/B7: even column (col%4==0), odd row. + barea = "B6/B7" + actualCb = i420aux.U[uvRow*i420aux.UStride+k2] + actualCr = i420aux.U[uvRow*i420aux.UStride+quarterW+k2] + } else { + // B8/B9: even column (col%4==2), odd row. + barea = "B8/B9" + actualCb = i420aux.V[uvRow*i420aux.VStride+k2] + actualCr = i420aux.V[uvRow*i420aux.VStride+quarterW+k2] + } + } + slog.Debug("H.264: pixel sample (LC=2 combine)", + "x", px, "y", py, + "area", barea, + "isIDR", stream2IsIDR, + "usedIDRSnapshot", yp == &g.avc444IDRYPlane, + "Y1", yp.data[py*yp.stride+px], + "Cb", actualCb, "Cr", actualCr, + "B", combined[off], "G", combined[off+1], "R", combined[off+2]) + } + if !g.lc2SampleLogged { + g.lc2SampleLogged = true + // B2/B3 (even col, even row) + lc2Sample(100, 50) + lc2Sample(500, 50) + // B4/B5 (odd col) — most important for diagnosing tint artifacts + lc2Sample(101, 50) + lc2Sample(501, 50) + lc2Sample(961, 50) + // B6/B7 (col%4==0, odd row) + lc2Sample(100, 51) + lc2Sample(500, 51) + // B8/B9 (col%4==2, odd row) + lc2Sample(102, 51) + lc2Sample(502, 51) + // video area — all four B-areas near the same spot + lc2Sample(960, 600) + lc2Sample(961, 600) + lc2Sample(960, 601) + lc2Sample(962, 601) + } else if !g.lc2PFrameSampleLogged && !stream2IsIDR { + g.lc2PFrameSampleLogged = true + lc2Sample(100, 50) + lc2Sample(101, 50) + lc2Sample(100, 51) + lc2Sample(102, 51) + lc2Sample(500, 50) + lc2Sample(501, 50) + lc2Sample(960, 400) + lc2Sample(961, 400) + lc2Sample(960, 401) + lc2Sample(962, 401) + lc2Sample(960, 600) + lc2Sample(961, 600) + } + decoded, pooled = cropBGRA(combined, w, h, destW, destH) + if w == destW && h == destH { + // cropBGRA returned combined unchanged; mark as pooled so caller releases it. + pooled = true + } else { + // cropBGRA created a new buffer; release the intermediate combined buffer. + releaseBitmapBuf(combined) + } + regions = stream2.regions + slog.Debug("RDPGFX: AVC444 LC=2 decoded", "w", w, "h", h, + "destW", destW, "destH", destH, "h264Len", len(stream2.h264Data)) + g.noteSuccessfulDecode() + return +} + +// softResetLimit is the number of in-place decoder recreations attempted +// before escalating to a full RDP reconnect. +const softResetLimit = 5 + +// maybeRequestKeyframe sends a keyframe request to the server when either +// decoder needs a fresh IDR. Requests are rate-limited to once per 2 seconds +// so that repeated nil-frame callbacks (e.g. while waiting for the IDR) don't +// flood the server. This covers both post-flush and post-soft-reset cases, +// including the case where h264dec2 was reset independently of h264dec. +// +// Proactive stall recovery: even when NeedsKeyframe()==false (decoder has not +// yet been reset), we send ForceRefresh early when the HW decoder appears to be +// stalling — packets are arriving but no real frame has been produced for longer +// than avc444YStaleness. This gives the server a ~1 second head-start to +// prepare an IDR before the stall detector fires and triggers SW fallback, +// reducing the visible freeze from ~18 s to a few seconds. +func (g *GfxHandler) maybeRequestKeyframe() { + if g.onKeyframeRequest == nil { + return + } + if g.h264dec == nil || g.h264dec.IsBroken() { + return + } + dec1NeedsKF := g.h264dec.NeedsKeyframe() + // Do NOT include h264dec2 here: ForceRefresh only triggers an LC=1 luma IDR + // from the server. The stream2/chroma IDR is never delivered via + // ForceRefresh — it arrives naturally as an LC=0 frame via primeAuxDecoder. + // Requesting ForceRefresh because h264dec2.NeedsIDR()=true spams the server + // with keyframe requests, causes the server to repeatedly send LC=1 IDRs, + // and can deadlock the main VideoToolbox decoder. + if !dec1NeedsKF { + // Proactive early request: if packets are flowing in but no real frame + // has been produced for avc444YStaleness, the HW decoder is likely + // producing null frames. Request a keyframe now so the server has time + // to respond before the stall detector escalates to SW fallback. + recvTime := g.h264dec.LastReceiveTime() + if recvTime.IsZero() || time.Since(recvTime) >= avc444YStaleness { + // No packets arriving — server is idle, not a HW stall. + return + } + lastNS := g.lastDecodedFrame.Load() + if lastNS == 0 || time.Since(time.Unix(0, lastNS)) < avc444YStaleness { + // Frames are still being produced recently — not stalling. + return + } + } + const keyframeRequestInterval = 2 * time.Second + if time.Since(g.lastKeyframeRequest) < keyframeRequestInterval { + return + } + g.lastKeyframeRequest = time.Now() + go g.onKeyframeRequest() +} + +// maybeNotifyDecoderBroken is called whenever the H.264 decoder returns a +// nil frame. It first tries up to softResetLimit in-place decoder resets +// (cheap: just recreate the FFmpeg/VideoToolbox context and ask the server +// for a fresh IDR). Only after all soft resets are exhausted does it call +// onDecoderBroken, which triggers a full RDP reconnect. +func (g *GfxHandler) maybeNotifyDecoderBroken() { + if g.decoderBrokenNotified { + return + } + if g.h264dec == nil || !g.h264dec.IsBroken() { + return + } + reason := g.h264dec.BrokenReason() + if reason == H264BrokenReasonNoIDR && g.h264dec.LastReceiveTime().IsZero() { + // The H.264 decoder's keyframe-wait timer fired, but the decoder has + // never received any data (LastReceiveTime is zero). This means the + // server is using a non-H.264 codec (e.g. CA Progressive / codecId=9) + // for the entire session and will never send H.264 frames. + // Sending ForceRefresh or reconnecting would disrupt the session + // unnecessarily — Ubuntu GNOME Remote Desktop responds to ForceRefresh + // with DEACTIVATEALLPDU followed by a disconnect. + // Disable the H.264 decoder so the watchdog can never fire again. + slog.Debug("H.264: watchdog fired but no H.264 data received — server uses non-H.264 codec, disabling H.264 decoder") + g.h264dec.Close() + g.h264dec = nil + return + } + if reason == H264BrokenReasonNoIDR { + // Allow one no-IDR soft reset before escalating to reconnect, unless + // we are already in SW fallback mode (after a HW stall). In the SW + // fallback case ForceRefresh was already sent multiple times during the + // VT stall and the server has not responded; another retry just prolongs + // the freeze by another keyframeWaitTimeoutSWFallback seconds. Skip + // straight to reconnect so the server can deliver a fresh IDR via the + // normal session-start path, which it reliably does. + // + // For the non-fallback path: ForceRefresh (SuppressOutput toggle) often + // fails to trigger a new AVC444 IDR from Windows servers; repeatedly + // retrying just prolongs the freeze. One attempt gives the server a + // fair chance; after that a full reconnect is faster. + // + // noIDRSoftResetCount is kept separate from softResetCount so that a + // prior HW-stall reset does not consume this budget — after an HW stall + // the SW fallback decoder skips retries (see above); for a pure SW + // session one no-IDR retry is still allowed. + const softResetLimitNoIDR = 1 + if !g.usingSWFallback && g.noIDRSoftResetCount < softResetLimitNoIDR { + g.noIDRSoftResetCount++ + slog.Debug("H.264: soft decoder reset (no-IDR)", + "attempt", g.noIDRSoftResetCount, "limit", softResetLimitNoIDR, + "reason", reason.String()) + g.h264dec.Close() + g.h264dec = newH264DecoderWithWatchdog(g.watchdogCh) + if g.h264dec2 != nil && g.h264dec2.IsBroken() { + slog.Debug("H.264: aux decoder also broken on soft reset, waiting for IDR to recreate") + g.h264dec2.Close() + g.h264dec2 = nil + } + g.lastKeyframeRequest = time.Time{} + g.maybeRequestKeyframe() + return + } + slog.Debug("H.264: escalating to reconnect after no-IDR soft reset exhausted", + "reason", reason.String()) + g.decoderBrokenNotified = true + if g.onDecoderBroken != nil { + go g.onDecoderBroken() + } + return + } + if g.softResetCount < softResetLimit { + g.softResetCount++ + if reason == H264BrokenReasonHWStall && !g.usingSWFallback { + // Switch to software (FFmpeg) decoding when VideoToolbox stalls. + // Even if a proactive ForceRefresh was already sent, the server + // typically delivers the IDR within ~1-2 s; the SW decoder will + // pick it up and the session continues without a full reconnect. + slog.Debug("H.264: HW stall — falling back to software decoding", + "attempt", g.softResetCount, "limit", softResetLimit) + g.usingSWFallback = true + } else { + slog.Debug("H.264: soft decoder reset", + "attempt", g.softResetCount, "limit", softResetLimit, + "reason", reason.String()) + } + g.h264dec.Close() + if g.usingSWFallback { + g.h264dec = newH264DecoderSWWithWatchdog(g.watchdogCh) + } else { + g.h264dec = newH264DecoderWithWatchdog(g.watchdogCh) + } + // Prime the SW fallback decoder with the last cached stream1 IDR so it + // can decode subsequent P-frames immediately, without waiting for the + // server to send a fresh IDR via ForceRefresh. + // + // We always prime when an IDR is cached, regardless of its age. A stale + // IDR is missing the reference frames decoded since then, so moving + // regions may show transient block noise until the next P-frames refresh + // them (or a fresh IDR fully heals the picture) — but for a mostly-static + // desktop the stale IDR is a close approximation and the artifacts are + // minor. Crucially this avoids the alternative: AVC444 servers only send + // an IDR at session start, so the cached IDR is essentially always "stale" + // at stall time; gating priming on freshness meant the SW decoder waited + // for a fresh IDR that never arrives, the watchdog fired, and the whole + // RDP session reconnected. Continuing with a primed SW decoder is far + // less disruptive than a reconnect. maybeRequestKeyframe() below still + // asks the server for a fresh IDR to clean up any residual artifacts. + idrAge := time.Since(g.lastStream1IDRTime) + idrFrameAge := g.framesDecoded.Load() - g.lastStream1IDRFrame + if g.usingSWFallback && len(g.lastStream1IDR) > 0 && + !g.lastStream1IDRTime.IsZero() { + slog.Debug("H.264: priming SW fallback with cached stream1 IDR to avoid IDR wait", + "idrLen", len(g.lastStream1IDR), + "idrAge", idrAge.Round(time.Millisecond), + "idrFrameAge", idrFrameAge, + ) + g.swFallbackPrimed = true + g.swFallbackDroppedCount = 0 + g.swFallbackFirstDropTime = time.Time{} + if _, err := g.h264dec.Decode(g.lastStream1IDR); err != nil { + slog.Debug("H.264: cached IDR prime failed, watchdog will wait for natural IDR", + "err", err) + } + } else if g.usingSWFallback { + slog.Debug("H.264: no cached stream1 IDR to prime SW fallback — watchdog will wait for a natural IDR", + "idrAge", idrAge.Round(time.Millisecond), + "idrFrameAge", idrFrameAge, + ) + } + // Keep h264dec2 if healthy; tear it down if already broken so + // primeAuxDecoder can recreate it when the next stream2 IDR arrives, + // rather than spinning up a new VT session only to have it break again. + // Always keep avc444YPlane so that LC=2 frames can continue to display + // stale-but-reasonable content during recovery. + if g.h264dec2 != nil && g.h264dec2.IsBroken() { + slog.Debug("H.264: aux decoder also broken on soft reset, waiting for IDR to recreate") + g.h264dec2.Close() + g.h264dec2 = nil + } + // Reset rate-limiter so keyframe request fires immediately after reset. + g.lastKeyframeRequest = time.Time{} + g.maybeRequestKeyframe() + return + } + // All soft resets exhausted — escalate to full reconnect. + g.decoderBrokenNotified = true + if g.onDecoderBroken != nil { + go g.onDecoderBroken() + } +} + +// swFallbackDropLimit is the minimum number of consecutive dropped frames, +// and swFallbackResyncTimeout the minimum elapsed time, that must accumulate +// after priming the SW fallback decoder with a stale cached IDR before we give +// up and reconnect. Both conditions must hold: a large frame count alone (a +// fast stall burst) should not reconnect before the ForceRefresh resync IDR has +// had a realistic chance to arrive, and a long idle gap alone should not +// reconnect if only one or two frames were bad. While corruption persists the +// frames are dropped (screen holds the last good frame) rather than shown, and +// maybeRequestKeyframe keeps asking the server for a fresh IDR; only when the +// server fails to heal within the timeout do we fall back to a full reconnect. +const swFallbackDropLimit = 3 +const swFallbackResyncTimeout = 2 * time.Second + +// trackSWFallbackDroppedFrame counts a dropped frame after a SW fallback IDR +// prime. When corruption from a stale prime persists past swFallbackDropLimit +// consecutive drops AND swFallbackResyncTimeout — i.e. the ForceRefresh resync +// IDR did not arrive in time — it marks the decoder broken so the application +// reconnects instead of showing a frozen/green screen. A genuine fresh IDR +// (maybeCacheStream1IDR) or any clean decode (noteSuccessfulDecode) clears the +// run before it reaches the escalation threshold. +func (g *GfxHandler) trackSWFallbackDroppedFrame() { + if !g.usingSWFallback || !g.swFallbackPrimed { + return + } + if g.swFallbackDroppedCount == 0 { + g.swFallbackFirstDropTime = time.Now() + } + g.swFallbackDroppedCount++ + if g.swFallbackDroppedCount >= swFallbackDropLimit && + !g.swFallbackFirstDropTime.IsZero() && + time.Since(g.swFallbackFirstDropTime) >= swFallbackResyncTimeout { + slog.Warn("H.264: SW fallback stale IDR prime did not resync in time, escalating to reconnect", + "dropped", g.swFallbackDroppedCount, + "persistedFor", time.Since(g.swFallbackFirstDropTime).Round(time.Millisecond)) + g.swFallbackPrimed = false + g.swFallbackDroppedCount = 0 + g.swFallbackFirstDropTime = time.Time{} + g.decoderBrokenNotified = true + if g.onDecoderBroken != nil { + go g.onDecoderBroken() + } + } +} + +// cropBGRA crops or pads BGRA pixel data to the target dimensions. +// When srcW == dstW and srcH == dstH the input slice is returned unchanged +// and pooled is false. Otherwise a new buffer is acquired from bitmapBufPool +// (pooled == true) and the caller must call releaseBitmapBuf on it. +func cropBGRA(src []byte, srcW, srcH, dstW, dstH int) ([]byte, bool) { + if srcW == dstW && srcH == dstH { + return src, false + } + out := acquireBitmapBuf(dstW * dstH * 4) + copyW := min(dstW, srcW) + copyH := min(dstH, srcH) + srcStride := srcW * 4 + dstStride := dstW * 4 + rowBytes := copyW * 4 + for y := range copyH { + copy(out[y*dstStride:y*dstStride+rowBytes], src[y*srcStride:y*srcStride+rowBytes]) + } + return out, true +} + +// avcRegionUseThresholdPercent is the upper bound on the *fraction* of the +// decoded frame area that the union of dirty rects can cover before we give +// up and just blit the whole frame. When the dirty area approaches the +// total area, the per-rect bookkeeping (allocation per rect, separate +// BitmapUpdate per rect) costs more than the bytes-copied savings. +const avcRegionUseThresholdPercent = 60 + +// shouldUseAVCRegions returns true when the per-region partial blit path is +// expected to be cheaper than a single full-frame blit. A single region +// covering everything is treated as "no win"; many tiny regions covering +// most of the frame are similarly bypassed. +func shouldUseAVCRegions(regions []avcRect, frameW, frameH int) bool { + if frameW <= 0 || frameH <= 0 { + return false + } + total := frameW * frameH + if total == 0 { + return false + } + // Sum (with overlap double-counting) — overlap is uncommon in practice + // and the threshold leaves slack for it. + sum := 0 + for _, r := range regions { + if r.right <= r.left || r.bottom <= r.top { + continue + } + w := int(r.right - r.left) + h := int(r.bottom - r.top) + sum += w * h + if sum*100 >= total*avcRegionUseThresholdPercent { + return false + } + } + return sum > 0 +} + +// blitAndEmitAVCRegions copies only the dirty rectangles of a decoded AVC +// frame into the persistent surface and emits a BitmapUpdate per region. +// All region coordinates are in decoded-frame space (i.e. relative to +// (left, top) on the surface). +// +// The emitted Data buffers are borrowed from bitmapBufPool and are returned +// to the pool once the synchronous onBitmap callback completes — see the +// BitmapUpdate lifecycle note. +func (g *GfxHandler) blitAndEmitAVCRegions(s *surface, left, top, frameW, frameH int, decoded []byte, regions []avcRect) { + frameStride := frameW * 4 + surfStride := int(s.width) * 4 + g.updatesBuf = g.updatesBuf[:0] + for _, rc := range regions { + if rc.right <= rc.left || rc.bottom <= rc.top { + continue + } + rx, ry := int(rc.left), int(rc.top) + rw, rh := int(rc.right-rc.left), int(rc.bottom-rc.top) + if rx+rw > frameW { + rw = frameW - rx + } + if ry+rh > frameH { + rh = frameH - ry + } + if rw <= 0 || rh <= 0 { + continue + } + rowBytes := rw * 4 + region := acquireBitmapBuf(rw * rh * 4) + for row := 0; row < rh; row++ { + srcOff := (ry+row)*frameStride + rx*4 + if srcOff+rowBytes > len(decoded) { + break + } + copy(region[row*rowBytes:row*rowBytes+rowBytes], + decoded[srcOff:srcOff+rowBytes]) + + // Mirror the same row into the persistent surface so any + // subsequent codec (RFX progressive etc.) operating on the + // same surface starts from the up-to-date pixels. + dy := top + ry + row + if dy < 0 || dy >= int(s.height) { + continue + } + dstOff := dy*surfStride + (left+rx)*4 + if dstOff < 0 || dstOff+rowBytes > len(s.data) { + continue + } + copy(s.data[dstOff:dstOff+rowBytes], + decoded[srcOff:srcOff+rowBytes]) + } + if !s.mapped || g.onBitmap == nil { + releaseBitmapBuf(region) + continue + } + destL := int(s.outputX) + left + rx + destT := int(s.outputY) + top + ry + g.updatesBuf = append(g.updatesBuf, BitmapUpdate{ + DestLeft: destL, DestTop: destT, + DestRight: destL + rw - 1, DestBottom: destT + rh - 1, + Width: rw, Height: rh, Bpp: 4, Data: region, + }) + } + g.emitAndReleaseUpdates(g.updatesBuf) +} + +// blitAVCRegionsToSurface copies only the dirty rectangles of a decoded AVC +// frame into the persistent CPU surface shadow (s.data), without allocating +// per-region buffers or emitting BitmapUpdates. It is the shadow-only +// counterpart to blitAndEmitAVCRegions, used on the GPU display path +// (onNV12/onI420) where the display is driven directly from the YUV planes and +// only the CPU shadow needs maintaining for later surface-to-surface / cache / +// mixed-codec operations. +// +// The caller must ensure the shadow is not stale (surface.shadowStale == false) +// before using this partial update; otherwise a full blitToSurface is required +// to repair regions that earlier GPU-only frames advanced without a shadow +// update. Region coordinates are in decoded-frame space (relative to +// (left, top) on the surface); both source and destination are bounds-clamped. +func (g *GfxHandler) blitAVCRegionsToSurface(s *surface, left, top, frameW, frameH int, decoded []byte, regions []avcRect) { + frameStride := frameW * 4 + surfStride := int(s.width) * 4 + for _, rc := range regions { + if rc.right <= rc.left || rc.bottom <= rc.top { + continue + } + rx, ry := int(rc.left), int(rc.top) + rw, rh := int(rc.right-rc.left), int(rc.bottom-rc.top) + if rx+rw > frameW { + rw = frameW - rx + } + if ry+rh > frameH { + rh = frameH - ry + } + if rw <= 0 || rh <= 0 { + continue + } + rowBytes := rw * 4 + for row := 0; row < rh; row++ { + srcOff := (ry+row)*frameStride + rx*4 + if srcOff+rowBytes > len(decoded) { + break + } + dy := top + ry + row + if dy < 0 || dy >= int(s.height) { + continue + } + dstOff := dy*surfStride + (left+rx)*4 + if dstOff < 0 || dstOff+rowBytes > len(s.data) { + continue + } + copy(s.data[dstOff:dstOff+rowBytes], decoded[srcOff:srcOff+rowBytes]) + } + } +} \ No newline at end of file diff --git a/plugin/rdpgfx/avc_test.go b/plugin/rdpgfx/avc_test.go new file mode 100644 index 0000000..da59350 --- /dev/null +++ b/plugin/rdpgfx/avc_test.go @@ -0,0 +1,33 @@ +package rdpgfx + +import ( + "encoding/binary" + "testing" +) + +func TestParseAVC420Stream(t *testing.T) { + data := make([]byte, 4+10+4) + binary.LittleEndian.PutUint32(data[:4], 1) + binary.LittleEndian.PutUint16(data[4:], 10) + binary.LittleEndian.PutUint16(data[6:], 20) + binary.LittleEndian.PutUint16(data[8:], 110) + binary.LittleEndian.PutUint16(data[10:], 220) + data[12] = 0x41 + data[13] = 0x7F + copy(data[14:], []byte{0x00, 0x00, 0x01, 0x65}) + + stream, err := parseAVC420Stream(data) + if err != nil { + t.Fatalf("parseAVC420Stream returned error: %v", err) + } + if len(stream.regions) != 1 { + t.Fatalf("expected 1 region, got %d", len(stream.regions)) + } + got := stream.regions[0] + if got.left != 10 || got.top != 20 || got.right != 110 || got.bottom != 220 { + t.Fatalf("unexpected region: %+v", got) + } + if string(stream.h264Data) != string([]byte{0x00, 0x00, 0x01, 0x65}) { + t.Fatalf("unexpected h264 payload: %v", stream.h264Data) + } +} diff --git a/plugin/rdpgfx/clear.go b/plugin/rdpgfx/clear.go new file mode 100644 index 0000000..0a56225 --- /dev/null +++ b/plugin/rdpgfx/clear.go @@ -0,0 +1,710 @@ +// clear.go 实现 MS-RDPEGFX 2.2.4 ClearCodec(TS_CLEARCODEC_BITMAP_STREAM) +// 的完整解码,算法对齐 FreeRDP libfreerdp/codec/clear.c。 +// +// 流结构:glyphFlags(1) + seqNumber(1) + [glyph 段] + residualByteCount(4) + +// bandsByteCount(4) + subcodecByteCount(4) + 三段载荷。 +package rdpgfx + +import ( + "encoding/binary" + "log/slog" +) + +const ( + clearFlagGlyphIndex = 0x01 + clearFlagGlyphHit = 0x02 + clearFlagCacheReset = 0x04 + + clearVBarSize = 32768 + clearShortVBarSize = 16384 + clearGlyphSize = 4000 +) + +// nsCodecDisabled 诊断开关:置 true 时丢弃 NSCodec 矩形 +var nsCodecDisabled = false + +type clearCodecCtx struct { + vBarStorage [clearVBarSize]vBarEntry + shortVBarStorage [clearShortVBarSize]vBarEntry + vBarCursor int + shortVBarCursor int + glyphCache [clearGlyphSize][]byte + seqNumber uint32 + logged [24]bool +} + +// logOnce 每类失败只打第一条日志,避免刷屏;位序号定位失败类别 +func (ctx *clearCodecCtx) logOnce(slot int, msg string, args ...any) { + if ctx.logged[slot] { + return + } + ctx.logged[slot] = true + slog.Warn("clear:"+msg, args...) +} + +func newClearCodecCtx() *clearCodecCtx { + return &clearCodecCtx{} +} + +// decode 解出一张 w×h 的 BGRA 位图;流非法时返回 nil——调用方必须保留 +// surface 原内容(对齐 FreeRDP 的整帧拒绝语义)。服务器会发送子编码矩形 +// 完全越界的"空更新"流,若不拒绝会把清零缓冲 blit 到屏幕形成花屏块。 +func (ctx *clearCodecCtx) decode(data []byte, w, h int) []byte { + out := make([]byte, w*h*4) + if len(data) < 2 { + return nil + } + off := 0 + glyphFlags := data[off] + seqNumber := data[off+1] + off += 2 + + // 序列号失步(丢包/上下文重置)时重新同步而不是丢弃整帧 + if ctx.seqNumber != uint32(seqNumber) { + ctx.seqNumber = uint32(seqNumber) + } + ctx.seqNumber = (ctx.seqNumber + 1) % 256 + + if glyphFlags&clearFlagCacheReset != 0 { + // FreeRDP 语义:重置游标并重新分配存储 → 所有缓存条目失效 + for i := range ctx.vBarStorage { + ctx.vBarStorage[i] = vBarEntry{} + } + for i := range ctx.shortVBarStorage { + ctx.shortVBarStorage[i] = vBarEntry{} + } + ctx.vBarCursor = 0 + ctx.shortVBarCursor = 0 + } + + // MS-RDPEGFX/FreeRDP 语义:GLYPH_HIT 必须与 GLYPH_INDEX 同置(0x03=命中), + // 单独的 HIT 是非法组合 + glyphIdx := -1 + if glyphFlags&clearFlagGlyphHit != 0 { + if glyphFlags&clearFlagGlyphIndex == 0 { + return nil + } + if off+2 > len(data) { + return nil + } + idx := int(binary.LittleEndian.Uint16(data[off:])) + off += 2 + if idx >= clearGlyphSize { + return nil + } + if cached := ctx.glyphCache[idx]; len(cached) == len(out) { + copy(out, cached) + } + return out + } + if glyphFlags&clearFlagGlyphIndex != 0 { + if off+2 > len(data) { + return nil + } + glyphIdx = int(binary.LittleEndian.Uint16(data[off:])) + off += 2 + if glyphIdx >= clearGlyphSize { + return nil + } + } + + if off+12 > len(data) { + // FreeRDP:GLYPH_HIT|INDEX 的纯命中流允许没有 12 字节长度头 + return nil + } + residualLen := int(binary.LittleEndian.Uint32(data[off:])) + bandsLen := int(binary.LittleEndian.Uint32(data[off+4:])) + subcodecLen := int(binary.LittleEndian.Uint32(data[off+8:])) + off += 12 + + if residualLen > 0 { + if off+residualLen > len(data) { + return nil + } + if !ctx.clearDecodeResidual(data[off:off+residualLen], w, h, out) { + return nil + } + } + off += residualLen + if bandsLen > 0 { + if off+bandsLen > len(data) { + return nil + } + if !ctx.clearDecodeBands(data[off:off+bandsLen], w, h, out) { + return nil + } + } + off += bandsLen + if subcodecLen > 0 { + if off+subcodecLen > len(data) { + return nil + } + if !ctx.clearDecodeSubcodecs(data[off:off+subcodecLen], w, h, out) { + return nil + } + } + + if glyphIdx >= 0 { + cached := make([]byte, len(out)) + copy(cached, out) + ctx.glyphCache[glyphIdx] = cached + } + return out +} + +// clearDecodeResidual:(b,g,r,runLen) 游程填充,runLen 用 1/2/4 字节变长编码。 +// 返回 false 表示流非法(对齐 FreeRDP 的整帧拒绝语义)。 +func (ctx *clearCodecCtx) clearDecodeResidual(data []byte, w, h int, out []byte) bool { + off, pixelIndex, pixelCount := 0, 0, w*h + for off < len(data) { + if off+4 > len(data) { + ctx.logOnce(2, "residual truncated (entry hdr)") + return false + } + b, g, r := data[off], data[off+1], data[off+2] + run := int(data[off+3]) + off += 4 + if run >= 0xFF { + if off+2 > len(data) { + ctx.logOnce(2, "residual truncated (run16 hdr)") + return false + } + run = int(binary.LittleEndian.Uint16(data[off:])) + off += 2 + if run >= 0xFFFF { + if off+4 > len(data) { + ctx.logOnce(2, "residual truncated (run32 hdr)") + return false + } + run = int(binary.LittleEndian.Uint32(data[off:])) + off += 4 + } + } + if run > pixelCount-pixelIndex { + ctx.logOnce(3, "residual run overflow", "run", run, "left", pixelCount-pixelIndex) + return false + } + i := pixelIndex * 4 + for n := 0; n < run; n++ { + out[i], out[i+1], out[i+2], out[i+3] = b, g, r, 0xFF + i += 4 + } + pixelIndex += run + } + if pixelIndex != pixelCount { + // FreeRDP:residual 必须恰好铺满整幅,否则整帧失败 + ctx.logOnce(16, "residual incomplete", "covered", pixelIndex, "want", pixelCount) + return false + } + return true +} + +// clearDecodeBands:vBar 列带解码(含短 vBar 缓存/整 vBar 缓存/背景合成)。 +// 返回 false 表示流非法(整帧拒绝)。 +func (ctx *clearCodecCtx) clearDecodeBands(data []byte, w, h int, out []byte) bool { + off := 0 + for off+11 <= len(data) { + xStart := int(binary.LittleEndian.Uint16(data[off:])) + xEnd := int(binary.LittleEndian.Uint16(data[off+2:])) + yStart := int(binary.LittleEndian.Uint16(data[off+4:])) + yEnd := int(binary.LittleEndian.Uint16(data[off+6:])) + cb, cg, cr := data[off+8], data[off+9], data[off+10] + off += 11 + if xEnd < xStart || yEnd < yStart { + ctx.logOnce(4, "bands bad rect", "xStart", xStart, "xEnd", xEnd, "yStart", yStart, "yEnd", yEnd) + return false + } + colorBkg := [4]byte{cb, cg, cr, 0xFF} + vBarCount := xEnd - xStart + 1 + vBarHeight := yEnd - yStart + 1 + if vBarHeight > 52 { + ctx.logOnce(5, "bands vBarHeight>52", "h", vBarHeight) + return false + } + + for i := 0; i < vBarCount; i++ { + if off+2 > len(data) { + return false + } + vBarHeader := binary.LittleEndian.Uint16(data[off:]) + off += 2 + + var entry *vBarEntry + var shortEntry *vBarEntry + vBarUpdate := false + vBarYOn := 0 + vBarShortPixelCount := 0 + + switch { + case vBarHeader&0xC000 == 0x4000: // SHORT_VBAR_CACHE_HIT + idx := int(vBarHeader & 0x3FFF) + if idx >= clearShortVBarSize { + ctx.logOnce(6, "bands short idx range", "idx", idx) + return false + } + shortEntry = &ctx.shortVBarStorage[idx] + if off >= len(data) { + return false + } + vBarYOn = int(data[off]) + off++ + vBarShortPixelCount = shortEntry.count + vBarUpdate = true + case vBarHeader&0xC000 == 0x0000: // SHORT_VBAR_CACHE_MISS + vBarYOn = int(vBarHeader & 0xFF) + vBarYOff := int((vBarHeader >> 8) & 0x3F) + if vBarYOff < vBarYOn { + ctx.logOnce(7, "bands yOff 52 { + ctx.logOnce(8, "bands short count>52", "n", vBarShortPixelCount) + return false + } + if off+vBarShortPixelCount*3 > len(data) { + return false + } + shortEntry = &ctx.shortVBarStorage[ctx.shortVBarCursor] + shortEntry.count = vBarShortPixelCount + shortEntry.pixels = make([]byte, vBarShortPixelCount*4) + for p := 0; p < vBarShortPixelCount; p++ { + shortEntry.pixels[p*4] = data[off+p*3] + shortEntry.pixels[p*4+1] = data[off+p*3+1] + shortEntry.pixels[p*4+2] = data[off+p*3+2] + shortEntry.pixels[p*4+3] = 0xFF + } + off += vBarShortPixelCount * 3 + ctx.shortVBarCursor = (ctx.shortVBarCursor + 1) % clearShortVBarSize + vBarUpdate = true + case vBarHeader&0x8000 == 0x8000: // VBAR_CACHE_HIT + idx := int(vBarHeader & 0x7FFF) + if idx >= clearVBarSize { + ctx.logOnce(9, "bands vbar idx range", "idx", idx) + return false + } + entry = &ctx.vBarStorage[idx] + if entry.pixels == nil { + // 缓存被重置后的命中:填充哑数据 + entry.count = vBarHeight + entry.pixels = make([]byte, vBarHeight*4) + } + default: + ctx.logOnce(10, "bands invalid vBarHeader", "hdr", vBarHeader) + return false + } + + if vBarUpdate { + ve := &ctx.vBarStorage[ctx.vBarCursor] + ve.count = vBarHeight + ve.pixels = make([]byte, vBarHeight*4) + // 前段背景 [0, vBarYOn) + bgFront := vBarYOn + if bgFront > vBarHeight { + bgFront = vBarHeight + } + for y2 := 0; y2 < bgFront; y2++ { + copy(ve.pixels[y2*4:(y2+1)*4], colorBkg[:]) + } + // 中段:短 vBar 像素 [vBarYOn, vBarYOn+vBarShortPixelCount) + count := vBarShortPixelCount + if vBarYOn+count > vBarHeight { + count = vBarHeight - vBarYOn + } + if count < 0 { + count = 0 + } + if count > 0 && shortEntry != nil { + copy(ve.pixels[vBarYOn*4:(vBarYOn+count)*4], shortEntry.pixels[0:count*4]) + } + // 后段背景 [vBarYOn+vBarShortPixelCount, vBarHeight) + start := vBarYOn + vBarShortPixelCount + if start < 0 { + start = 0 + } + if start > vBarHeight { + start = vBarHeight + } + for y2 := start; y2 < vBarHeight; y2++ { + copy(ve.pixels[y2*4:(y2+1)*4], colorBkg[:]) + } + ctx.vBarCursor = (ctx.vBarCursor + 1) % clearVBarSize + entry = ve + } + + if entry == nil || entry.count != vBarHeight { + continue + } + + // 落到输出:第 i 列,从 yStart 起取 min(count, h-yStart) 像素 + x := xStart + i + if x >= w { + continue + } + count := entry.count + if count > h-yStart { + count = h - yStart + } + if count <= 0 { + continue + } + for y := 0; y < count; y++ { + dst := (yStart+y)*w*4 + x*4 + copy(out[dst:dst+4], entry.pixels[y*4:y*4+4]) + } + } + } + return true +} + +// clearDecodeSubcodecs:矩形级子编码(0=BGR24 原始,1=NSCodec,2=RLEX)。 +// 返回 false 表示流非法(整帧拒绝)。 +func (ctx *clearCodecCtx) clearDecodeSubcodecs(data []byte, w, h int, out []byte) bool { + off := 0 + for off+13 <= len(data) { + xStart := int(binary.LittleEndian.Uint16(data[off:])) + yStart := int(binary.LittleEndian.Uint16(data[off+2:])) + width := int(binary.LittleEndian.Uint16(data[off+4:])) + height := int(binary.LittleEndian.Uint16(data[off+6:])) + bmpLen := int(binary.LittleEndian.Uint32(data[off+8:])) + subcodecId := data[off+12] + off += 13 + if off+bmpLen > len(data) { + ctx.logOnce(11, "subcodec payload truncated", "id", subcodecId, "want", bmpLen, "have", len(data)-off) + return false + } + payload := data[off : off+bmpLen] + off += bmpLen + + // FreeRDP 语义:子编码矩形越界是非法流,整帧拒绝(否则服务器发来的 + // 完全越界"空更新"矩形会把清零缓冲 blit 到屏幕形成花屏块) + if xStart >= w || yStart >= h || xStart+width > w || yStart+height > h { + ctx.logOnce(17, "subcodec rect out of bounds", + "id", subcodecId, "x", xStart, "y", yStart, "rw", width, "rh", height, "w", w, "h", h) + return false + } + + switch subcodecId { + case 0: // 未压缩 BGR24 + if len(payload) != width*height*3 { + ctx.logOnce(18, "subcodec BGR24 size mismatch", "want", width*height*3, "have", len(payload)) + return false + } + clearWriteBGR24(payload, width, height, out, xStart, yStart, w, h) + case 1: // NSCodec(YCoCg 色域压缩编码) + if !clearDecodeNSCodec(payload, width, height, out, xStart, yStart, w, h) { + return false + } + case 2: // RLEX + if !ctx.clearDecodeRLEX(payload, width, height, out, xStart, yStart, w, h) { + return false + } + default: + ctx.logOnce(15, "unknown subcodec id", "id", subcodecId) + return false + } + } + return true +} + +func clearWriteBGR24(data []byte, w, h int, out []byte, xDst, yDst, surfW, surfH int) { + stride := w * 3 + if stride*h > len(data) { + return + } + for y := 0; y < h; y++ { + if yDst+y >= surfH { + return + } + row := data[y*stride : (y+1)*stride] + for x := 0; x < w; x++ { + if xDst+x >= surfW { + break + } + dst := (yDst+y)*surfW*4 + (xDst+x)*4 + out[dst], out[dst+1], out[dst+2], out[dst+3] = row[x*3], row[x*3+1], row[x*3+2], 0xFF + } + } +} + +// clearDecodeNSCodec 解码 NSCodec 子编码矩形,算法对齐 FreeRDP libfreerdp/codec/nsc.c: +// 20 字节头(4×PlaneByteCount + ColorLossLevel + ChromaSubsamplingLevel + 保留 2 字节), +// 4 个色度平面(Y/Co/Cg/A)RLE 解压后做色损恢复与 YCoCg→RGB 转换。 + +func clearDecodeNSCodec(data []byte, w, h int, out []byte, xDst, yDst, surfW, surfH int) bool { + if len(data) < 20 { + return false + } + var planeCount [4]int + total := 0 + for i := 0; i < 4; i++ { + planeCount[i] = int(binary.LittleEndian.Uint32(data[i*4:])) + total += planeCount[i] + } + colorLossLevel := int(data[16]) + chroma := int(data[17]) + if colorLossLevel < 1 || colorLossLevel > 7 { + return false + } + planes := data[20:] + if len(planes) < total { + return false + } + + rw := (w + 7) &^ 7 + rh := (h + 1) &^ 1 + // 原始平面字节数(OrgByteCount) + org := [4]int{w * h, w * h, w * h, w * h} + if chroma != 0 { + org[0] = rw * h + org[1] = (rw / 2) * (rh / 2) + org[2] = org[1] + } + var pbuf [4][]byte + for i := 0; i < 4; i++ { + pbuf[i] = make([]byte, org[i]) + } + + off := 0 + for i := 0; i < 4; i++ { + psize := planeCount[i] + plane := planes[off : off+psize] + off += psize + switch { + case psize == 0: + for j := range pbuf[i] { + pbuf[i][j] = 0xFF + } + case psize < org[i]: + if !nscRLEDecode(plane, pbuf[i]) { + return false + } + default: + copy(pbuf[i], plane[:org[i]]) + } + } + + // 色损恢复 + YCoCg→RGB:shift = ColorLossLevel-1 + shift := uint(colorLossLevel - 1) + bmp := make([]byte, w*h*4) + pos := 0 + for y := 0; y < h; y++ { + var yplane, coplane, cgplane []byte + if chroma != 0 { + yplane = pbuf[0][y*rw:] + coplane = pbuf[1][(y>>1)*(rw>>1):] + cgplane = pbuf[2][(y>>1)*(rw>>1):] + } else { + yplane = pbuf[0][y*w:] + coplane = pbuf[1][y*w:] + cgplane = pbuf[2][y*w:] + } + // A 平面(pbuf[3])不参与显示:canvas putImageData 会保留 alpha, + // 非 255 的 alpha 会呈现半透明色块,故强制输出不透明 + coIdx, cgIdx := 0, 0 + for x := 0; x < w; x++ { + yv := int16(yplane[x]) + // C: (INT16)(INT8)(((INT16)u8) << shift) —— 16 位环绕后截断低 8 位再符号扩展 + cov := int16(int8(byte(int16(coplane[coIdx]) << shift))) + cgv := int16(int8(byte(int16(cgplane[cgIdx]) << shift))) + rv := yv + cov - cgv + gv := yv + cgv + bv := yv - cov - cgv + bmp[pos] = nscClamp(bv) + bmp[pos+1] = nscClamp(gv) + bmp[pos+2] = nscClamp(rv) + // A 平面值不作为透明度使用:RDP 桌面内容恒不透明,而 canvas + // putImageData 会保留 alpha,若 A 平面含非 255 值会呈现半透明色块 + bmp[pos+3] = 0xFF + pos += 4 + if chroma != 0 { + if x%2 == 1 { + coIdx++ + cgIdx++ + } + } else { + coIdx++ + cgIdx++ + } + } + } + blitBGRA(bmp, w, h, out, xDst, yDst, surfW, surfH) + return true +} + +// nscRLEDecode 单平面 NSC 游程解码,对齐 FreeRDP nsc_rle_decode +func nscRLEDecode(in, out []byte) bool { + inOff, outOff := 0, 0 + left := len(out) + for left > 4 { + if inOff >= len(in) { + return false + } + value := in[inOff] + inOff++ + if left == 5 { + out[outOff] = value + outOff++ + left-- + } else if inOff >= len(in) { + return false + } else if value == in[inOff] { + inOff++ + if inOff >= len(in) { + return false + } + var run int + if in[inOff] < 0xFF { + run = int(in[inOff]) + 2 + inOff++ + } else { + if inOff+5 > len(in) { + return false + } + inOff++ + run = int(in[inOff]) | int(in[inOff+1])<<8 | int(in[inOff+2])<<16 | int(in[inOff+3])<<24 + inOff += 4 + } + if run > len(out)-outOff || run > left { + return false + } + for j := 0; j < run; j++ { + out[outOff+j] = value + } + outOff += run + left -= run + } else { + out[outOff] = value + outOff++ + left-- + } + } + if len(out)-outOff < 4 || left < 4 || len(in)-inOff < 4 { + return false + } + copy(out[outOff:outOff+4], in[inOff:inOff+4]) + return true +} + +func nscClamp(v int16) byte { + if v < 0 { + return 0 + } + if v > 255 { + return 255 + } + return byte(v) +} + +func blitBGRA(src []byte, w, h int, out []byte, xDst, yDst, surfW, surfH int) { + for y := 0; y < h; y++ { + dy := yDst + y + if dy >= surfH { + return + } + for x := 0; x < w; x++ { + dx := xDst + x + if dx >= surfW { + break + } + copy(out[(dy*surfW+dx)*4:(dy*surfW+dx)*4+4], src[(y*w+x)*4:(y*w+x)*4+4]) + } + } +} + +// clearDecodeRLEX:调色板游程编码(套位打包索引) +func (ctx *clearCodecCtx) clearDecodeRLEX(data []byte, w, h int, out []byte, xDst, yDst, surfW, surfH int) bool { + if len(data) < 1 { + return false + } + paletteCount := int(data[0]) + if paletteCount < 1 || paletteCount > 127 { + return false + } + if 1+paletteCount*3 > len(data) { + return false + } + palette := make([][4]byte, paletteCount) + off := 1 + for i := 0; i < paletteCount; i++ { + palette[i] = [4]byte{data[off], data[off+1], data[off+2], 0xFF} // b,g,r + off += 3 + } + + numBits := clearLog2Floor(paletteCount-1) + 1 + pixelCount := w * h + pixelIndex := 0 + x, y := 0, 0 + putPixel := func(c [4]byte) { + if xDst+x < surfW && yDst+y < surfH { + dst := (yDst+y)*surfW*4 + (xDst+x)*4 + copy(out[dst:dst+4], c[:]) + } + if x++; x >= w { + y++ + x = 0 + } + } + for off+2 <= len(data) && pixelIndex < pixelCount { + tmp := data[off] + run := int(data[off+1]) + off += 2 + suiteDepth := int(tmp >> uint(numBits) & clear8BitMask(8-numBits)) + stopIndex := int(tmp & clear8BitMask(numBits)) + startIndex := stopIndex - suiteDepth + if run >= 0xFF { + if off+2 > len(data) { + ctx.logOnce(13, "rlex truncated (run16 hdr)") + return false + } + run = int(binary.LittleEndian.Uint16(data[off:])) + off += 2 + if run >= 0xFFFF { + if off+4 > len(data) { + ctx.logOnce(13, "rlex truncated (run32 hdr)") + return false + } + run = int(binary.LittleEndian.Uint32(data[off:])) + off += 4 + } + } + if startIndex < 0 || startIndex >= paletteCount || stopIndex >= paletteCount { + ctx.logOnce(12, "rlex bad palette index", "start", startIndex, "stop", stopIndex, "count", paletteCount) + return false + } + if run > pixelCount-pixelIndex { + ctx.logOnce(14, "rlex run overflow", "run", run, "left", pixelCount-pixelIndex) + return false + } + for i := 0; i < run; i++ { + putPixel(palette[startIndex]) + } + pixelIndex += run + for i := 0; i <= suiteDepth && pixelIndex < pixelCount; i++ { + putPixel(palette[startIndex+i]) + pixelIndex++ + } + } + if pixelIndex != pixelCount { + // FreeRDP:RLEX 必须恰好覆盖矩形,否则整帧失败 + ctx.logOnce(19, "rlex incomplete", "covered", pixelIndex, "want", pixelCount) + return false + } + return true +} + +func clear8BitMask(bits int) byte { + if bits <= 0 || bits > 8 { + return 0 + } + return byte((1 << uint(bits)) - 1) +} + +func clearLog2Floor(v int) int { + log := 0 + for v > 1 { + v >>= 1 + log++ + } + return log +} diff --git a/plugin/rdpgfx/clear_test.go b/plugin/rdpgfx/clear_test.go new file mode 100644 index 0000000..94eccb1 --- /dev/null +++ b/plugin/rdpgfx/clear_test.go @@ -0,0 +1,295 @@ +package rdpgfx + +// ClearCodec 解码器单元测试:按 MS-RDPEGFX 2.2.4 逐段手工构造码流, +// 验证 residual 游程、bands 列带(含短 vBar 缓存命中)、RLEX 子编码与 glyph 缓存。 + +import ( + "encoding/binary" + "testing" +) + +// buildClearStream 组装一个完整 TS_CLEARCODEC_BITMAP_STREAM +func buildClearStream(glyphFlags, seqNumber byte, glyphIndex uint16, + residual, bands, subcodec []byte) []byte { + out := []byte{glyphFlags, seqNumber} + if glyphFlags&(clearFlagGlyphHit|clearFlagGlyphIndex) != 0 { + out = binary.LittleEndian.AppendUint16(out, glyphIndex) + } + out = binary.LittleEndian.AppendUint32(out, uint32(len(residual))) + out = binary.LittleEndian.AppendUint32(out, uint32(len(bands))) + out = binary.LittleEndian.AppendUint32(out, uint32(len(subcodec))) + out = append(out, residual...) + out = append(out, bands...) + out = append(out, subcodec...) + return out +} + +// px 取出输出位图中 (x,y) 的 BGRA +func px(out []byte, w, x, y int) [4]byte { + i := (y*w + x) * 4 + return [4]byte{out[i], out[i+1], out[i+2], out[i+3]} +} + +func expectPixel(t *testing.T, out []byte, w, x, y int, want [4]byte, ctx string) { + t.Helper() + got := px(out, w, x, y) + if got != want { + t.Fatalf("%s: pixel(%d,%d) = %v, want %v", ctx, x, y, got, want) + } +} + +// TestClearResidual:单条游程铺满整幅 +func TestClearResidual(t *testing.T) { + w, h := 4, 2 + // residual: b,g,r=0x10,0x20,0x30 runLen=8 + residual := []byte{0x10, 0x20, 0x30, 8} + stream := buildClearStream(0, 0, 0, residual, nil, nil) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + for y := 0; y < h; y++ { + for x := 0; x < w; x++ { + expectPixel(t, out, w, x, y, [4]byte{0x10, 0x20, 0x30, 0xFF}, "residual") + } + } +} + +// TestClearResidualRunEncoding:长游程 2 字节/4 字节变长编码 +func TestClearResidualRunEncoding(t *testing.T) { + w, h := 16, 16 // 256 像素 + // run=255(触发 2 字节编码 0x00FF)+ run=1 → 共 256 + residual := []byte{1, 2, 3, 0xFF, 0xFF, 0x00, 4, 5, 6, 1} + stream := buildClearStream(0, 0, 0, residual, nil, nil) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + for y := 0; y < 16; y++ { + for x := 0; x < 16; x++ { + want := [4]byte{1, 2, 3, 0xFF} + if y == 15 && x == 15 { + want = [4]byte{4, 5, 6, 0xFF} + } + expectPixel(t, out, w, x, y, want, "residual-run") + } + } +} + +// TestClearBandsShortVBarCacheMiss:列带 + 短 vBar 存缓存 + 背景合成 +func TestClearBandsShortVBarCacheMiss(t *testing.T) { + w, h := 2, 4 + // band: x0..1, y0..3 → vBarCount=2, vBarHeight=4 + bands := []byte{} + bands = binary.LittleEndian.AppendUint16(bands, 0) // xStart + bands = binary.LittleEndian.AppendUint16(bands, 1) // xEnd + bands = binary.LittleEndian.AppendUint16(bands, 0) // yStart + bands = binary.LittleEndian.AppendUint16(bands, 3) // yEnd + bands = append(bands, 0xAA, 0xBB, 0xCC) // cb,cg,cr 背景色 + // 每列:SHORT_VBAR_CACHE_MISS,yOn=1,yOff=3 → 2 个像素,位于行 1-2 + for col := 0; col < 2; col++ { + header := uint16(1) | uint16(3)<<8 // yOn=1, yOff=3 + bands = binary.LittleEndian.AppendUint16(bands, header) + bands = append(bands, byte(col+1), 0x00, 0x10) // 行1: b,g,r + bands = append(bands, byte(col+2), 0x00, 0x20) // 行2 + } + stream := buildClearStream(0, 0, 0, nil, bands, nil) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + for col := 0; col < 2; col++ { + expectPixel(t, out, w, col, 0, [4]byte{0xAA, 0xBB, 0xCC, 0xFF}, "bands-bg-front") + expectPixel(t, out, w, col, 1, [4]byte{byte(col + 1), 0x00, 0x10, 0xFF}, "bands-short-1") + expectPixel(t, out, w, col, 2, [4]byte{byte(col + 2), 0x00, 0x20, 0xFF}, "bands-short-2") + expectPixel(t, out, w, col, 3, [4]byte{0xAA, 0xBB, 0xCC, 0xFF}, "bands-bg-back") + } +} + +// TestClearBandsShortVBarCacheHit:第二帧复用第一帧的短 vBar 缓存 +func TestClearBandsShortVBarCacheHit(t *testing.T) { + w, h := 1, 4 + mkBands := func(hit bool) []byte { + bands := []byte{} + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 3) + bands = append(bands, 0x01, 0x02, 0x03) // 背景 + if hit { + // SHORT_VBAR_CACHE_HIT:idx=0 + 1 字节 yOn=2 + bands = binary.LittleEndian.AppendUint16(bands, 0x4000|0) + bands = append(bands, 2) + } else { + // MISS:yOn=2,yOff=4 → 2 像素 + bands = binary.LittleEndian.AppendUint16(bands, 2|4<<8) + bands = append(bands, 0x10, 0x20, 0x30) + bands = append(bands, 0x40, 0x50, 0x60) + } + return bands + } + ctx := newClearCodecCtx() + ctx.decode(buildClearStream(0, 0, 0, nil, mkBands(false), nil), w, h) + out := ctx.decode(buildClearStream(0, 1, 0, nil, mkBands(true), nil), w, h) + // 缓存命中:行 2-3 = 缓存的两像素,行 0-1 = 背景 + expectPixel(t, out, w, 0, 0, [4]byte{0x01, 0x02, 0x03, 0xFF}, "hit-bg-0") + expectPixel(t, out, w, 0, 1, [4]byte{0x01, 0x02, 0x03, 0xFF}, "hit-bg-1") + expectPixel(t, out, w, 0, 2, [4]byte{0x10, 0x20, 0x30, 0xFF}, "hit-short-2") + expectPixel(t, out, w, 0, 3, [4]byte{0x40, 0x50, 0x60, 0xFF}, "hit-short-3") +} + +// TestClearSubcodecBGR24:未压缩 BGR24 子编码 +func TestClearSubcodecBGR24(t *testing.T) { + w, h := 2, 2 + sub := []byte{} + sub = binary.LittleEndian.AppendUint16(sub, 1) // xStart + sub = binary.LittleEndian.AppendUint16(sub, 1) // yStart + sub = binary.LittleEndian.AppendUint16(sub, 1) // width + sub = binary.LittleEndian.AppendUint16(sub, 1) // height + sub = binary.LittleEndian.AppendUint32(sub, 3) // bitmapDataByteCount + sub = append(sub, 0) // subcodecId = BGR24 + sub = append(sub, 0x77, 0x88, 0x99) // b,g,r + stream := buildClearStream(0, 0, 0, nil, nil, sub) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + expectPixel(t, out, w, 0, 0, [4]byte{0, 0, 0, 0}, "bgr24-untouched") + expectPixel(t, out, w, 1, 1, [4]byte{0x77, 0x88, 0x99, 0xFF}, "bgr24-hit") +} + +// TestClearRLEX:调色板游程子编码(套位打包索引) +func TestClearRLEX(t *testing.T) { + w, h := 4, 1 + sub := []byte{} + sub = binary.LittleEndian.AppendUint16(sub, 0) // xStart + sub = binary.LittleEndian.AppendUint16(sub, 0) // yStart + sub = binary.LittleEndian.AppendUint16(sub, 4) // width + sub = binary.LittleEndian.AppendUint16(sub, 1) // height + sub = binary.LittleEndian.AppendUint32(sub, 15) + // RLEX: subcodecId=2 + paletteCount=4 + 4×(b,g,r) + sub = append(sub, 2) + sub = append(sub, 4) + sub = append(sub, 0x10, 0x00, 0x00) + sub = append(sub, 0x20, 0x00, 0x00) + sub = append(sub, 0x30, 0x00, 0x00) + sub = append(sub, 0x40, 0x00, 0x00) + // paletteCount=4 → numBits=2;suiteDepth=3, stopIndex=3 → tmp=(3<<2)|3=0x0F + // runLen=0(1 字节)→ 输出 4 像素 palette[0..3] + sub = append(sub, 0x0F, 0) + stream := buildClearStream(0, 0, 0, nil, nil, sub) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + expectPixel(t, out, w, 0, 0, [4]byte{0x10, 0, 0, 0xFF}, "rlex-0") + expectPixel(t, out, w, 1, 0, [4]byte{0x20, 0, 0, 0xFF}, "rlex-1") + expectPixel(t, out, w, 2, 0, [4]byte{0x30, 0, 0, 0xFF}, "rlex-2") + expectPixel(t, out, w, 3, 0, [4]byte{0x40, 0, 0, 0xFF}, "rlex-3") +} + +// TestClearGlyphIndexThenHit:GLYPH_INDEX 缓存输出,GLYPH_HIT 直接取缓存 +func TestClearGlyphIndexThenHit(t *testing.T) { + w, h := 2, 2 + residual := []byte{0x11, 0x22, 0x33, 4} + ctx := newClearCodecCtx() + // 第一帧:GLYPH_INDEX,缓存到 index=7 + ctx.decode(buildClearStream(clearFlagGlyphIndex, 0, 7, residual, nil, nil), w, h) + // 第二帧:GLYPH_HIT|GLYPH_INDEX(合法命中组合),无载荷也应输出缓存 + out := ctx.decode(buildClearStream(clearFlagGlyphHit|clearFlagGlyphIndex, 1, 7, nil, nil, nil), w, h) + for y := 0; y < h; y++ { + for x := 0; x < w; x++ { + expectPixel(t, out, w, x, y, [4]byte{0x11, 0x22, 0x33, 0xFF}, "glyph-hit") + } + } + // 单独的 GLYPH_HIT 是非法组合 → 整帧拒绝(nil,调用方保留 surface) + out = ctx.decode(buildClearStream(clearFlagGlyphHit, 2, 7, nil, nil, nil), w, h) + if out != nil { + t.Fatalf("invalid glyph flags: want nil frame, got %d bytes", len(out)) + } + // 空更新流(子编码矩形完全越界,服务器真实样本形态)→ 整帧拒绝 + sub := []byte{} + sub = binary.LittleEndian.AppendUint16(sub, 32) // xStart == w(越界) + sub = binary.LittleEndian.AppendUint16(sub, 64) // yStart == h(越界) + sub = binary.LittleEndian.AppendUint16(sub, 8) // width + sub = binary.LittleEndian.AppendUint16(sub, 0) // height + sub = binary.LittleEndian.AppendUint32(sub, 15) + sub = append(sub, 2) // subcodecId = RLEX + sub = append(sub, 2, 2, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0) + out = ctx.decode(buildClearStream(0, 3, 0, nil, nil, sub), 32, 64) + if out != nil { + t.Fatalf("out-of-bounds subcodec rect: want nil frame, got %d bytes", len(out)) + } +} + +// TestClearVBarCacheHitAcrossFrames:整 vBar 缓存跨帧复用(CACHE_RESET 后哑填充不崩溃) +func TestClearVBarCacheHitAcrossFrames(t *testing.T) { + w, h := 1, 2 + mkBands := func(hit bool) []byte { + bands := []byte{} + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 0) + bands = binary.LittleEndian.AppendUint16(bands, 1) // yEnd → vBarHeight=2 + bands = append(bands, 0, 0, 0) + if hit { + bands = binary.LittleEndian.AppendUint16(bands, 0x8000|0) // VBAR_CACHE_HIT idx 0 + } else { + // MISS(vBarHeader 最高两位为 00 且 0x8000 未置位):count=yOn..yOff 不可用, + // 直接用像素数:yOn=0, yOff=2 → header = 0|2<<8 + bands = binary.LittleEndian.AppendUint16(bands, 0|2<<8) + bands = append(bands, 0x0A, 0x0B, 0x0C, 0x0D, 0x0E, 0x0F) + } + return bands + } + ctx := newClearCodecCtx() + ctx.decode(buildClearStream(0, 0, 0, nil, mkBands(false), nil), w, h) + // 带 CACHE_RESET 重置 vBar 存储,再用 idx0 命中 → 哑填充不 panic + out := ctx.decode(buildClearStream(clearFlagCacheReset, 1, 0, nil, mkBands(true), nil), w, h) + expectPixel(t, out, w, 0, 0, [4]byte{0, 0, 0, 0}, "vbar-reset-dummy") +} + +// TestClearNSCodec:NSCodec 子编码(RLE 平面 + YCoCg 恢复 + 空 alpha 平面) +func TestClearNSCodec(t *testing.T) { + w, h := 8, 2 // org 每平面 16 字节 + // 单平面 RLE:12 字节游程 [v,v,run-2] + 尾部 4 字节原样 = 7 字节 + rle := func(v byte) []byte { return []byte{v, v, 10, v, v, v, v} } + yPlane := rle(0x80) + coPlane := rle(0x40) + cgPlane := rle(0x00) + + sub := []byte{} + sub = binary.LittleEndian.AppendUint16(sub, 0) // xStart + sub = binary.LittleEndian.AppendUint16(sub, 0) // yStart + sub = binary.LittleEndian.AppendUint16(sub, uint16(w)) // width + sub = binary.LittleEndian.AppendUint16(sub, uint16(h)) // height + sub = binary.LittleEndian.AppendUint32(sub, uint32(20 + 7 + 7 + 7 + 0)) + sub = append(sub, 1) // subcodecId = NSCodec + for _, n := range []int{7, 7, 7, 0} { + sub = binary.LittleEndian.AppendUint32(sub, uint32(n)) // PlaneByteCount + } + sub = append(sub, 1, 0, 0, 0) // ColorLossLevel=1, ChromaSubsampling=0, 保留 + sub = append(sub, yPlane...) + sub = append(sub, coPlane...) + sub = append(sub, cgPlane...) + // A 平面 byteCount=0 → 全 0xFF + + stream := buildClearStream(0, 0, 0, nil, nil, sub) + ctx := newClearCodecCtx() + out := ctx.decode(stream, w, h) + // shift=0:r=0x80+0x40-0x00=0xC0, g=0x80, b=0x80-0x40-0x00=0x40 + for y := 0; y < h; y++ { + for x := 0; x < w; x++ { + expectPixel(t, out, w, x, y, [4]byte{0x40, 0x80, 0xC0, 0xFF}, "nsc") + } + } +} + +// TestProgDWTExtrapolateConstant:常数 LL3 + 零高频 → IDWT 输出应为同一常数 +func TestProgDWTExtrapolateConstant(t *testing.T) { + buf := make([]int16, 4096) + // extrapolate 布局 LL3 在 [4015,4096),9×9 行主序 + for y := 0; y < 9; y++ { + for x := 0; x < 9; x++ { + buf[4015+y*9+x] = 128 + } + } + progDWTExtrapolate(buf) + for y := 0; y < 64; y++ { + for x := 0; x < 64; x++ { + if got := buf[y*64+x]; got != 128 { + t.Fatalf("DWT(%d,%d) = %d, want 128", x, y, got) + } + } + } +} diff --git a/plugin/rdpgfx/convert_parallel_test.go b/plugin/rdpgfx/convert_parallel_test.go new file mode 100644 index 0000000..b0d2dc7 --- /dev/null +++ b/plugin/rdpgfx/convert_parallel_test.go @@ -0,0 +1,229 @@ +package rdpgfx + +import ( + "math/rand" + "testing" +) + +// TestParallelRowsCoverage verifies parallelRows invokes fn over contiguous, +// non-overlapping chunks that together cover [0,h) exactly once, for both the +// serial (small) and parallel (large) regimes. +func TestParallelRowsCoverage(t *testing.T) { + for _, dim := range []struct{ w, h int }{ + {16, 16}, // tiny → serial + {64, 64}, // small → serial + {256, 256}, // at threshold → parallel + {512, 300}, // parallel, height not divisible by worker count + {1920, 1080}, // typical full-screen → parallel + {8, 1}, // single row + {8, 0}, // zero rows + } { + counts := make([]int32, dim.h) + parallelRows(dim.w, dim.h, func(y0, y1 int) { + for r := y0; r < y1; r++ { + counts[r]++ + } + }) + for r := 0; r < dim.h; r++ { + if counts[r] != 1 { + t.Fatalf("dim %dx%d: row %d visited %d times, want 1", dim.w, dim.h, r, counts[r]) + } + } + } +} + +// serialI420ToBGRA is an independent reference implementation used to verify the +// parallelised i420ToBGRA output is identical regardless of how rows are split. +func serialI420ToBGRA(src *H264FrameI420) []byte { + w, h := src.Width, src.Height + out := make([]byte, w*h*4) + for row := 0; row < h; row++ { + yOff := row * src.YStride + uOff := (row >> 1) * src.UStride + vOff := (row >> 1) * src.VStride + for col := 0; col < w; col++ { + uv := col >> 1 + o := (row*w + col) * 4 + u := int(src.U[uOff+uv]) - 128 + v := int(src.V[vOff+uv]) - 128 + if src.FullRange { + y := int(src.Y[yOff+col]) + out[o] = clampByte((256*y + 475*u + 128) >> 8) + out[o+1] = clampByte((256*y - 48*u - 120*v + 128) >> 8) + out[o+2] = clampByte((256*y + 403*v + 128) >> 8) + } else { + c := int(src.Y[yOff+col]) - 16 + out[o] = clampByte((298*c + 541*u + 128) >> 8) + out[o+1] = clampByte((298*c - 55*u - 136*v + 128) >> 8) + out[o+2] = clampByte((298*c + 459*v + 128) >> 8) + } + out[o+3] = 255 + } + } + return out +} + +func makeI420(w, h int, fullRange bool, rng *rand.Rand) *H264FrameI420 { + cw, ch := (w+1)/2, (h+1)/2 + f := &H264FrameI420{ + Y: make([]byte, w*h), U: make([]byte, cw*ch), V: make([]byte, cw*ch), + YStride: w, UStride: cw, VStride: cw, + Width: w, Height: h, FullRange: fullRange, + } + for i := range f.Y { + f.Y[i] = byte(rng.Intn(256)) + } + for i := range f.U { + f.U[i] = byte(rng.Intn(256)) + f.V[i] = byte(rng.Intn(256)) + } + return f +} + +// TestI420ToBGRAParallelMatchesSerial checks that the parallelised conversion +// produces byte-identical output to the reference for both small (serial) and +// large (parallel) frames, in full- and limited-range modes. +func TestI420ToBGRAParallelMatchesSerial(t *testing.T) { + rng := rand.New(rand.NewSource(1)) + for _, dim := range []struct{ w, h int }{ + {64, 64}, // serial path + {640, 480}, // parallel path + {512, 511}, // parallel, odd height + } { + for _, fr := range []bool{false, true} { + f := makeI420(dim.w, dim.h, fr, rng) + got, pooled := i420ToBGRA(f) + if got == nil { + t.Fatalf("i420ToBGRA returned nil for %dx%d", dim.w, dim.h) + } + want := serialI420ToBGRA(f) + if len(got) != len(want) { + t.Fatalf("%dx%d fr=%v: len %d != %d", dim.w, dim.h, fr, len(got), len(want)) + } + for i := range want { + if got[i] != want[i] { + t.Fatalf("%dx%d fr=%v: byte %d = %d, want %d", dim.w, dim.h, fr, i, got[i], want[i]) + } + } + if pooled { + releaseBitmapBuf(got) + } + } + } +} + +func BenchmarkI420ToBGRA1080p(b *testing.B) { +rng := rand.New(rand.NewSource(2)) +f := makeI420(1920, 1080, false, rng) +b.SetBytes(int64(1920 * 1080 * 4)) +b.ResetTimer() +for i := 0; i < b.N; i++ { +out, pooled := i420ToBGRA(f) +if pooled { +releaseBitmapBuf(out) +} +} +} + +// serialCombineAVC444v2BGRA mirrors combineAVC444v2BGRA's exact indexing in a +// single serial loop, used to confirm the parallelised version is byte-identical +// regardless of how the rows are split across workers. +func serialCombineAVC444v2BGRA(yPlane []byte, yStride int, cachedU, cachedV []byte, uvStride int, +i420aux *H264FrameI420, fullRange bool, w, h int) []byte { +out := make([]byte, w*h*4) +halfW := w / 2 +quarterW := w / 4 +for row := 0; row < h; row++ { +yRowOff := row * yStride +uvRow := row >> 1 +uvRowOff := uvRow * uvStride +auxYRowOff := row * i420aux.YStride +auxURowOff := uvRow * i420aux.UStride +auxVRowOff := uvRow * i420aux.VStride +outIdx := row * w * 4 +for col := 0; col < w; col++ { +Y := yPlane[yRowOff+col] +var Cb, Cr byte +if col&1 == 1 { +k := col >> 1 +Cb = i420aux.Y[auxYRowOff+k] +Cr = i420aux.Y[auxYRowOff+halfW+k] +} else if row&1 == 0 { +k := col >> 1 +Cb = cachedU[uvRowOff+k] +Cr = cachedV[uvRowOff+k] +} else { +k := col >> 2 +if col&2 == 0 { +Cb = i420aux.U[auxURowOff+k] +Cr = i420aux.U[auxURowOff+quarterW+k] +} else { +Cb = i420aux.V[auxVRowOff+k] +Cr = i420aux.V[auxVRowOff+quarterW+k] +} +} +u := int(Cb) - 128 +v := int(Cr) - 128 +if fullRange { +y := int(Y) +out[outIdx] = clampByte((256*y + 475*u + 128) >> 8) +out[outIdx+1] = clampByte((256*y - 48*u - 120*v + 128) >> 8) +out[outIdx+2] = clampByte((256*y + 403*v + 128) >> 8) +} else { +c := int(Y) - 16 +out[outIdx] = clampByte((298*c + 541*u + 128) >> 8) +out[outIdx+1] = clampByte((298*c - 55*u - 136*v + 128) >> 8) +out[outIdx+2] = clampByte((298*c + 459*v + 128) >> 8) +} +out[outIdx+3] = 255 +outIdx += 4 +} +} +return out +} + +func TestCombineAVC444v2BGRAParallelMatchesSerial(t *testing.T) { +rng := rand.New(rand.NewSource(3)) +for _, dim := range []struct{ w, h int }{ +{64, 64}, // serial path +{640, 480}, // parallel path +{512, 510}, // parallel, even dims +} { +w, h := dim.w, dim.h +uvStride := (w + 1) / 2 +uvH := (h + 1) / 2 +yPlane := make([]byte, w*h) +cachedU := make([]byte, uvStride*uvH) +cachedV := make([]byte, uvStride*uvH) +for i := range yPlane { +yPlane[i] = byte(rng.Intn(256)) +} +for i := range cachedU { +cachedU[i] = byte(rng.Intn(256)) +cachedV[i] = byte(rng.Intn(256)) +} +aux := &H264FrameI420{ +Y: make([]byte, w*h), U: make([]byte, uvStride*uvH), V: make([]byte, uvStride*uvH), +YStride: w, UStride: uvStride, VStride: uvStride, Width: w, Height: h, +} +for i := range aux.Y { +aux.Y[i] = byte(rng.Intn(256)) +} +for i := range aux.U { +aux.U[i] = byte(rng.Intn(256)) +aux.V[i] = byte(rng.Intn(256)) +} +for _, fr := range []bool{false, true} { +got, pooled := combineAVC444v2BGRA(yPlane, w, cachedU, cachedV, uvStride, aux, fr, w, h, nil) +want := serialCombineAVC444v2BGRA(yPlane, w, cachedU, cachedV, uvStride, aux, fr, w, h) +for i := range want { +if got[i] != want[i] { +t.Fatalf("%dx%d fr=%v: byte %d = %d, want %d", w, h, fr, i, got[i], want[i]) +} +} +if pooled { +releaseBitmapBuf(got) +} +} +} +} diff --git a/plugin/rdpgfx/ffmpeg/h264_ffmpeg.go b/plugin/rdpgfx/ffmpeg/h264_ffmpeg.go new file mode 100644 index 0000000..1523cd2 --- /dev/null +++ b/plugin/rdpgfx/ffmpeg/h264_ffmpeg.go @@ -0,0 +1,3077 @@ +//go:build h264 + +package ffmpeg + +/* +#cgo pkg-config: libavcodec libavutil libswscale +#cgo nocallback avcodec_alloc_context3 +#cgo nocallback avcodec_find_decoder +#cgo nocallback avcodec_flush_buffers +#cgo nocallback avcodec_free_context +#cgo nocallback avcodec_get_hw_config +#cgo nocallback avcodec_open2 +#cgo nocallback avcodec_receive_frame +#cgo nocallback avcodec_send_packet +#cgo nocallback av_buffer_ref +#cgo nocallback av_buffer_unref +#cgo nocallback av_frame_alloc +#cgo nocallback av_frame_free +#cgo nocallback av_frame_unref +#cgo nocallback av_hwdevice_ctx_create +#cgo nocallback av_hwdevice_get_type_name +#cgo nocallback av_hwdevice_iterate_types +#cgo nocallback av_hwframe_transfer_data +#cgo nocallback av_packet_alloc +#cgo nocallback av_packet_free +#cgo nocallback grdp_find_v4l2m2m +#cgo nocallback grdp_hwframe_map +#cgo nocallback grdp_is_full_range_fmt +#cgo nocallback grdp_set_get_format +#cgo nocallback grdp_set_hw_pix_fmt +#cgo nocallback grdp_set_low_delay +#cgo nocallback grdp_suppress_av_log +#cgo nocallback grdp_yuvj_to_yuv +#cgo nocallback sws_freeContext +#cgo nocallback sws_getContext +#cgo nocallback grdp_sws_set_src_range +#cgo noescape avcodec_send_packet +#cgo noescape grdp_copy_yuv420p_to_i420 +#cgo nocallback grdp_copy_yuv420p_to_i420 +#cgo noescape grdp_copy_nv12_to_i420 +#cgo nocallback grdp_copy_nv12_to_i420 +#cgo noescape grdp_copy_nv12 +#cgo nocallback grdp_copy_nv12 +#cgo noescape grdp_yuv420p_to_bgra_regions +#cgo nocallback grdp_yuv420p_to_bgra_regions +#cgo noescape grdp_yuv420p_to_bgra_rows +#cgo nocallback grdp_yuv420p_to_bgra_rows +#cgo noescape grdp_nv12_to_bgra_regions +#cgo nocallback grdp_nv12_to_bgra_regions +#cgo noescape grdp_nv12_to_bgra_rows +#cgo nocallback grdp_nv12_to_bgra_rows +#cgo noescape grdp_frame_to_bgra +#cgo nocallback grdp_frame_to_bgra +#cgo noescape grdp_sample_nv12 +#cgo nocallback grdp_sample_nv12 +#cgo noescape grdp_sample_yuv +#cgo nocallback grdp_sample_yuv +#cgo noescape grdp_sample_nv12_at +#cgo nocallback grdp_sample_nv12_at +#cgo nocallback grdp_is_warmup_nv12 +#cgo noescape grdp_is_low_chroma_nv12 +#cgo nocallback grdp_is_low_chroma_nv12 +#cgo noescape grdp_is_low_chroma_yuv420p +#cgo nocallback grdp_is_low_chroma_yuv420p +#include +#include +#include +#include +#include +#include +#include +#include +#ifdef __ARM_NEON__ +#include +#endif +#ifdef __SSE2__ +#include +#endif + +// grdp_suppress_av_log sets FFmpeg's global log level to FATAL so that +// decoder-level error messages (e.g. "sps_id out of range", "no frame!") +// are not printed to stderr. Those messages are expected and harmless +// during H.264 stream recovery; grdp emits its own slog warnings instead. +static void grdp_suppress_av_log(void) { + av_log_set_level(AV_LOG_FATAL); +} + +// get_format callback that prefers the hardware pixel format stored in opaque. +static enum AVPixelFormat grdp_get_hw_format( + AVCodecContext *ctx, const enum AVPixelFormat *pix_fmts) { + enum AVPixelFormat hw_fmt = (enum AVPixelFormat)(intptr_t)ctx->opaque; + if (hw_fmt == AV_PIX_FMT_NONE) return pix_fmts[0]; + for (const enum AVPixelFormat *p = pix_fmts; *p != AV_PIX_FMT_NONE; p++) { + if (*p == hw_fmt) return *p; + } + return pix_fmts[0]; +} + +static void grdp_set_get_format(AVCodecContext *ctx) { + ctx->get_format = grdp_get_hw_format; +} + +// grdp_set_low_delay enables AV_CODEC_FLAG_LOW_DELAY on the codec context +// so the decoder emits frames as soon as they are decoded, without waiting +// to reorder B-frames. RDP H.264 streams transmit in display order and do +// not use B-frame reordering, so the default reorder buffer only adds +// apparent latency and (on VideoToolbox) makes legitimate frames look like +// "null frames" to our stall detector, triggering spurious hard resets. +static void grdp_set_low_delay(AVCodecContext *ctx) { + ctx->flags |= AV_CODEC_FLAG_LOW_DELAY; + ctx->flags2 |= AV_CODEC_FLAG2_FAST; +} + +static void grdp_set_hw_pix_fmt(AVCodecContext *ctx, enum AVPixelFormat fmt) { + ctx->opaque = (void*)(intptr_t)fmt; +} + +// grdp_hwframe_map attempts a zero-copy CPU mapping of a hardware frame. +// On VideoToolbox (macOS), decoded frames live in IOSurface-backed memory that +// is accessible from both CPU and GPU. av_hwframe_map creates a mapped view +// without copying the pixel data, allowing NV12 extraction without the extra +// GPU→RAM copy that av_hwframe_transfer_data would perform. +// Returns 0 on success; callers must fall back to av_hwframe_transfer_data +// on negative return (hardware type does not support mapping). +static int grdp_hwframe_map(AVFrame *dst, const AVFrame *src) { + return av_hwframe_map(dst, src, AV_HWFRAME_MAP_READ); +} + +// Helper: convert AVFrame to BGRA via swscale. +static int grdp_frame_to_bgra(struct SwsContext *sws, + AVFrame *src, uint8_t *dst, int dst_stride) { + uint8_t *dst_data[4] = {dst, NULL, NULL, NULL}; + int dst_linesize[4] = {dst_stride, 0, 0, 0}; + return sws_scale(sws, + (const uint8_t *const *)src->data, src->linesize, + 0, src->height, + dst_data, dst_linesize); +} + +// Map deprecated YUVJ pixel formats to their non-J equivalents. +// YUVJ formats are full-range YUV; the modern way is to use the plain YUV +// format and communicate the range via sws_setColorspaceDetails. +static enum AVPixelFormat grdp_yuvj_to_yuv(enum AVPixelFormat fmt) { + switch (fmt) { + case AV_PIX_FMT_YUVJ420P: return AV_PIX_FMT_YUV420P; + case AV_PIX_FMT_YUVJ422P: return AV_PIX_FMT_YUV422P; + case AV_PIX_FMT_YUVJ444P: return AV_PIX_FMT_YUV444P; + case AV_PIX_FMT_YUVJ440P: return AV_PIX_FMT_YUV440P; + default: return fmt; + } +} + +// Return 1 if fmt is a full-range (YUVJ) format, 0 otherwise. +static int grdp_is_full_range_fmt(enum AVPixelFormat fmt) { + return (fmt == AV_PIX_FMT_YUVJ420P || + fmt == AV_PIX_FMT_YUVJ422P || + fmt == AV_PIX_FMT_YUVJ444P || + fmt == AV_PIX_FMT_YUVJ440P) ? 1 : 0; +} + +// grdp_bt601_pixel writes one BGRA pixel using BT.601 coefficients. +// u and v are pre-offset (i.e. raw_value - 128). +// full_range: 0 = limited (video) range [16-235 / 16-240], +// 1 = full range [0-255]. +#define CLAMP8(x) ((x) < 0 ? 0 : (x) > 255 ? 255 : (uint8_t)(x)) +static inline void grdp_bt601_pixel( + int y_raw, int u, int v, int full_range, uint8_t *dst) +{ + int r, g, b; + if (full_range) { + int y = y_raw; + r = (256*y + 359*v + 128) >> 8; + g = (256*y - 88*u - 183*v + 128) >> 8; + b = (256*y + 454*u + 128) >> 8; + } else { + int c = y_raw - 16; + r = (298*c + 409*v + 128) >> 8; + g = (298*c - 100*u - 208*v + 128) >> 8; + b = (298*c + 516*u + 128) >> 8; + } + dst[0] = CLAMP8(b); + dst[1] = CLAMP8(g); + dst[2] = CLAMP8(r); + dst[3] = 255; +} + +// grdp_yuv420p_to_bgra converts a planar YUV420P/YUVJ420P frame to packed +// BGRA using BT.601 coefficients. This bypasses swscale entirely so that +// the broken ARM64 colorspace-matrix fallback path is never taken. +#ifdef __ARM_NEON__ +// grdp_yuv420p_to_bgra_neon_8 processes 8 luma pixels (4 UV pairs) per call. +// For YUV420P each UV sample covers 2 horizontal luma pixels; we load 4 U and +// 4 V bytes and duplicate each with vzip to produce 8 per-pixel U/V vectors, +// then follow the same NEON arithmetic path as grdp_nv12_to_bgra_neon_8. +static inline void grdp_yuv420p_to_bgra_neon_8( + const uint8_t *yrow, const uint8_t *urow, const uint8_t *vrow, + uint8_t *drow, int col, + int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, + int16_t yoff) +{ + // Load 8 luma bytes, convert to int16, subtract Y offset (16 or 0). + uint8x8_t y_u8 = vld1_u8(yrow + col); + int16x8_t c16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(y_u8)), + vdupq_n_s16(yoff)); + + // Load 4 U and 4 V bytes (one UV pair per 2 luma pixels). + // vzip duplicates each byte: [U0,U1,U2,U3,...] → [U0,U0,U1,U1,U2,U2,U3,U3]. + // ffmpeg pads AVFrame line buffers for SIMD so loading 8 bytes is safe. + uint8x8_t u_raw = vld1_u8(urow + (col >> 1)); + uint8x8_t v_raw = vld1_u8(vrow + (col >> 1)); + uint8x8_t u8 = vzip_u8(u_raw, u_raw).val[0]; + uint8x8_t v8 = vzip_u8(v_raw, v_raw).val[0]; + + int16x8_t u16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u8)), vdupq_n_s16(128)); + int16x8_t v16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v8)), vdupq_n_s16(128)); + + // Compute R/G/B with int32 to avoid overflow. Process 4+4 pixels. + int16x4_t c_lo = vget_low_s16(c16), u_lo = vget_low_s16(u16), v_lo = vget_low_s16(v16); + int16x4_t c_hi = vget_high_s16(c16), u_hi = vget_high_s16(u16), v_hi = vget_high_s16(v16); + + int32x4_t ky_lo = vmull_n_s16(c_lo, ky), ky_hi = vmull_n_s16(c_hi, ky); + + int32x4_t r_lo = vaddq_s32(vaddq_s32(ky_lo, vmull_n_s16(v_lo, kr)), vdupq_n_s32(128)); + int32x4_t g_lo = vaddq_s32(vsubq_s32(vsubq_s32(ky_lo, vmull_n_s16(u_lo, kgu)), vmull_n_s16(v_lo, kgv)), vdupq_n_s32(128)); + int32x4_t b_lo = vaddq_s32(vaddq_s32(ky_lo, vmull_n_s16(u_lo, kb)), vdupq_n_s32(128)); + + int32x4_t r_hi = vaddq_s32(vaddq_s32(ky_hi, vmull_n_s16(v_hi, kr)), vdupq_n_s32(128)); + int32x4_t g_hi = vaddq_s32(vsubq_s32(vsubq_s32(ky_hi, vmull_n_s16(u_hi, kgu)), vmull_n_s16(v_hi, kgv)), vdupq_n_s32(128)); + int32x4_t b_hi = vaddq_s32(vaddq_s32(ky_hi, vmull_n_s16(u_hi, kb)), vdupq_n_s32(128)); + + // Shift >>8, saturate int32→int16→uint8, store interleaved BGRA. + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(r_lo,8)), vqmovn_s32(vshrq_n_s32(r_hi,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(g_lo,8)), vqmovn_s32(vshrq_n_s32(g_hi,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(b_lo,8)), vqmovn_s32(vshrq_n_s32(b_hi,8)))); + uint8x8x4_t bgra; + bgra.val[0] = b; + bgra.val[1] = g; + bgra.val[2] = r; + bgra.val[3] = vdup_n_u8(255); + vst4_u8(drow + col * 4, bgra); +} + +// grdp_yuv420p_to_bgra_neon_16 processes 16 luma pixels (8 UV pairs) per call. +// One vld1_u8 covers all 8 U (and V) samples needed for 16 luma columns. +// vzip_u8 duplicates the low half for pixels 0-7 and the high half for 8-15, +// giving twice the throughput of grdp_yuv420p_to_bgra_neon_8 per iteration. +static inline void grdp_yuv420p_to_bgra_neon_16( + const uint8_t *yrow, const uint8_t *urow, const uint8_t *vrow, + uint8_t *drow, int col, + int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, + int16_t yoff) +{ + // Load 16 luma bytes, subtract Y offset. + uint8x16_t y_u8 = vld1q_u8(yrow + col); + int16x8_t c_lo = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(vget_low_u8(y_u8))), + vdupq_n_s16(yoff)); + int16x8_t c_hi = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(vget_high_u8(y_u8))), + vdupq_n_s16(yoff)); + + // Load 8 U and 8 V bytes; duplicate each to cover its two luma pixels. + uint8x8_t u_raw = vld1_u8(urow + (col >> 1)); + uint8x8_t v_raw = vld1_u8(vrow + (col >> 1)); + uint8x8x2_t u_zip = vzip_u8(u_raw, u_raw); // val[0]=[U0,U0,...,U3,U3] val[1]=[U4,U4,...,U7,U7] + uint8x8x2_t v_zip = vzip_u8(v_raw, v_raw); + + // Pixels 0-7 (low half). + { + int16x8_t u8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u_zip.val[0])), vdupq_n_s16(128)); + int16x8_t v8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v_zip.val[0])), vdupq_n_s16(128)); + int16x4_t cl = vget_low_s16(c_lo), ul = vget_low_s16(u8), vl = vget_low_s16(v8); + int16x4_t ch = vget_high_s16(c_lo), uh = vget_high_s16(u8), vh = vget_high_s16(v8); + int32x4_t kyl = vmull_n_s16(cl, ky), kyh = vmull_n_s16(ch, ky); + int32x4_t rl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(vl, kr)), vdupq_n_s32(128)); + int32x4_t gl = vaddq_s32(vsubq_s32(vsubq_s32(kyl, vmull_n_s16(ul, kgu)), vmull_n_s16(vl, kgv)), vdupq_n_s32(128)); + int32x4_t bl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(ul, kb)), vdupq_n_s32(128)); + int32x4_t rh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(vh, kr)), vdupq_n_s32(128)); + int32x4_t gh = vaddq_s32(vsubq_s32(vsubq_s32(kyh, vmull_n_s16(uh, kgu)), vmull_n_s16(vh, kgv)), vdupq_n_s32(128)); + int32x4_t bh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(uh, kb)), vdupq_n_s32(128)); + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(rl,8)), vqmovn_s32(vshrq_n_s32(rh,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(gl,8)), vqmovn_s32(vshrq_n_s32(gh,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(bl,8)), vqmovn_s32(vshrq_n_s32(bh,8)))); + uint8x8x4_t bgra; bgra.val[0]=b; bgra.val[1]=g; bgra.val[2]=r; bgra.val[3]=vdup_n_u8(255); + vst4_u8(drow + col * 4, bgra); + } + // Pixels 8-15 (high half). + { + int16x8_t u8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u_zip.val[1])), vdupq_n_s16(128)); + int16x8_t v8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v_zip.val[1])), vdupq_n_s16(128)); + int16x4_t cl = vget_low_s16(c_hi), ul = vget_low_s16(u8), vl = vget_low_s16(v8); + int16x4_t ch = vget_high_s16(c_hi), uh = vget_high_s16(u8), vh = vget_high_s16(v8); + int32x4_t kyl = vmull_n_s16(cl, ky), kyh = vmull_n_s16(ch, ky); + int32x4_t rl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(vl, kr)), vdupq_n_s32(128)); + int32x4_t gl = vaddq_s32(vsubq_s32(vsubq_s32(kyl, vmull_n_s16(ul, kgu)), vmull_n_s16(vl, kgv)), vdupq_n_s32(128)); + int32x4_t bl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(ul, kb)), vdupq_n_s32(128)); + int32x4_t rh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(vh, kr)), vdupq_n_s32(128)); + int32x4_t gh = vaddq_s32(vsubq_s32(vsubq_s32(kyh, vmull_n_s16(uh, kgu)), vmull_n_s16(vh, kgv)), vdupq_n_s32(128)); + int32x4_t bh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(uh, kb)), vdupq_n_s32(128)); + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(rl,8)), vqmovn_s32(vshrq_n_s32(rh,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(gl,8)), vqmovn_s32(vshrq_n_s32(gh,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(bl,8)), vqmovn_s32(vshrq_n_s32(bh,8)))); + uint8x8x4_t bgra; bgra.val[0]=b; bgra.val[1]=g; bgra.val[2]=r; bgra.val[3]=vdup_n_u8(255); + vst4_u8(drow + (col + 8) * 4, bgra); + } +} +#endif // __ARM_NEON__ + +// ---------------------------------------------------------------------------- +// SSE2 YUV→BGRA inline helpers (x86_64) +// Each function converts 8 luma pixels per call using 128-bit SIMD. +// Coefficients and arithmetic match the NEON paths above (BT.601, fixed-point +// with 8 fractional bits): R = (ky*Y + kr*V + 128) >> 8, etc. +// _mm_madd_epi16 is used for the R and B channels because kb=516 overflows +// int16 multiplication (516×127 = 65532 > 32767); madd widens to int32. +// ---------------------------------------------------------------------------- +#ifdef __SSE2__ + +// Shared BT.601 arithmetic core. Receives pre-biased int16 vectors +// y16 (Y − yoff), u16 (U − 128), v16 (V − 128) and stores 8 BGRA pixels. +static inline void grdp_yuv_to_bgra_sse2_core( + __m128i y16, __m128i u16, __m128i v16, + uint8_t *drow, int col, + int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb) +{ + const __m128i add128 = _mm_set1_epi32(128); + const __m128i alpha = _mm_set1_epi8((char)255); + // Interleaved coefficient pairs for _mm_madd_epi16: elem0=ky, elem1=k? + const __m128i coeff_r = _mm_set_epi16(kr, ky, kr, ky, + kr, ky, kr, ky); + const __m128i coeff_b = _mm_set_epi16(kb, ky, kb, ky, + kb, ky, kb, ky); + // Green uses ky*Y − kgv*V (via madd) then adds −kgu*U (sign-extended). + const __m128i coeff_yv = _mm_set_epi16((int16_t)-kgv, ky, (int16_t)-kgv, ky, + (int16_t)-kgv, ky, (int16_t)-kgv, ky); + const __m128i neg_kgu = _mm_set1_epi16((int16_t)-kgu); + + __m128i yv_lo = _mm_unpacklo_epi16(y16, v16); + __m128i yu_lo = _mm_unpacklo_epi16(y16, u16); + __m128i r32_lo = _mm_add_epi32(_mm_madd_epi16(yv_lo, coeff_r), add128); + __m128i b32_lo = _mm_add_epi32(_mm_madd_epi16(yu_lo, coeff_b), add128); + __m128i g32_lo = _mm_madd_epi16(yv_lo, coeff_yv); + { + __m128i ngu = _mm_mullo_epi16(u16, neg_kgu); + // Sign-extend int16 ngu to int32 using srai-15 trick, then unpack. + g32_lo = _mm_add_epi32(g32_lo, + _mm_add_epi32(_mm_unpacklo_epi16(ngu, _mm_srai_epi16(ngu, 15)), add128)); + } + __m128i yv_hi = _mm_unpackhi_epi16(y16, v16); + __m128i yu_hi = _mm_unpackhi_epi16(y16, u16); + __m128i r32_hi = _mm_add_epi32(_mm_madd_epi16(yv_hi, coeff_r), add128); + __m128i b32_hi = _mm_add_epi32(_mm_madd_epi16(yu_hi, coeff_b), add128); + __m128i g32_hi = _mm_madd_epi16(yv_hi, coeff_yv); + { + __m128i ngu = _mm_mullo_epi16(u16, neg_kgu); + g32_hi = _mm_add_epi32(g32_hi, + _mm_add_epi32(_mm_unpackhi_epi16(ngu, _mm_srai_epi16(ngu, 15)), add128)); + } + + // Shift right 8 (un-scale), pack int32→int16→uint8 (auto-saturating clamp). + __m128i r16 = _mm_packs_epi32(_mm_srai_epi32(r32_lo, 8), _mm_srai_epi32(r32_hi, 8)); + __m128i g16 = _mm_packs_epi32(_mm_srai_epi32(g32_lo, 8), _mm_srai_epi32(g32_hi, 8)); + __m128i b16 = _mm_packs_epi32(_mm_srai_epi32(b32_lo, 8), _mm_srai_epi32(b32_hi, 8)); + __m128i r8 = _mm_packus_epi16(r16, r16); + __m128i g8 = _mm_packus_epi16(g16, g16); + __m128i b8 = _mm_packus_epi16(b16, b16); + + // Interleave B,G,R,A into two 16-byte stores covering 8 BGRA pixels. + __m128i bg = _mm_unpacklo_epi8(b8, g8); + __m128i ra = _mm_unpacklo_epi8(r8, alpha); + _mm_storeu_si128((__m128i *)(drow + col * 4), _mm_unpacklo_epi16(bg, ra)); + _mm_storeu_si128((__m128i *)(drow + col * 4 + 16), _mm_unpackhi_epi16(bg, ra)); +} + +// 8-pixel NV12 (semi-planar Y + interleaved UV) → BGRA. +static inline void grdp_nv12_to_bgra_sse2_8( + const uint8_t *yrow, const uint8_t *uvrow, uint8_t *drow, int col, + int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, int16_t yoff) +{ + const __m128i zero = _mm_setzero_si128(); + __m128i y8 = _mm_loadl_epi64((const __m128i *)(yrow + col)); + __m128i y16 = _mm_sub_epi16(_mm_unpacklo_epi8(y8, zero), _mm_set1_epi16(yoff)); + + // Load 8 interleaved UV bytes [U0,V0,U1,V1,...,U3,V3]. + __m128i uv8 = _mm_loadl_epi64((const __m128i *)(uvrow + col)); + // Isolate U (even bytes) and V (odd bytes), pack to low 8 bytes. + __m128i u_pack = _mm_packus_epi16(_mm_and_si128(uv8, _mm_set1_epi16((int16_t)0x00FF)), zero); + __m128i v_pack = _mm_packus_epi16(_mm_srli_epi16(uv8, 8), zero); + // Duplicate each sample to cover its two luma pixels, then extend to int16. + __m128i u16 = _mm_sub_epi16( + _mm_unpacklo_epi8(_mm_unpacklo_epi8(u_pack, u_pack), zero), _mm_set1_epi16(128)); + __m128i v16 = _mm_sub_epi16( + _mm_unpacklo_epi8(_mm_unpacklo_epi8(v_pack, v_pack), zero), _mm_set1_epi16(128)); + + grdp_yuv_to_bgra_sse2_core(y16, u16, v16, drow, col, ky, kr, kgu, kgv, kb); +} + +// 8-pixel planar YUV420P → BGRA. +// FFmpeg guarantees at least 64-byte padding at end of each line buffer, so +// loading 8 bytes when only 4 U/V bytes are needed per 8 luma pixels is safe. +static inline void grdp_yuv420p_to_bgra_sse2_8( + const uint8_t *yrow, const uint8_t *urow, const uint8_t *vrow, + uint8_t *drow, int col, + int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, int16_t yoff) +{ + const __m128i zero = _mm_setzero_si128(); + __m128i y8 = _mm_loadl_epi64((const __m128i *)(yrow + col)); + __m128i y16 = _mm_sub_epi16(_mm_unpacklo_epi8(y8, zero), _mm_set1_epi16(yoff)); + + // Load 4 U and 4 V bytes (8-byte load; only low 4 used after duplication). + __m128i u4 = _mm_loadl_epi64((const __m128i *)(urow + (col >> 1))); + __m128i u16 = _mm_sub_epi16( + _mm_unpacklo_epi8(_mm_unpacklo_epi8(u4, u4), zero), _mm_set1_epi16(128)); + __m128i v4 = _mm_loadl_epi64((const __m128i *)(vrow + (col >> 1))); + __m128i v16 = _mm_sub_epi16( + _mm_unpacklo_epi8(_mm_unpacklo_epi8(v4, v4), zero), _mm_set1_epi16(128)); + + grdp_yuv_to_bgra_sse2_core(y16, u16, v16, drow, col, ky, kr, kgu, kgv, kb); +} + +#endif // __SSE2__ + +static void grdp_yuv420p_to_bgra( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range) +{ + int width = src->width; + int height = src->height; +#ifdef __ARM_NEON__ + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 15 < width; col += 16) + grdp_yuv420p_to_bgra_neon_16(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < width; col += 8) + grdp_yuv420p_to_bgra_neon_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + // Scalar tail for widths not a multiple of 8. + for (; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 7 < width; col += 8) + grdp_yuv420p_to_bgra_sse2_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#else + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + for (int col = 0; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#endif +} + +// grdp_yuv420p_to_bgra_regions is the region-aware variant of +// grdp_yuv420p_to_bgra. Only pixels within the n_rects dirty rectangles +// (flat array of [left,top,right,bottom] uint16 tuples) are written to dst; +// all other pixels are left untouched, saving work proportional to the +// fraction of the frame that did not change. +static void grdp_yuv420p_to_bgra_regions( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range, + const uint16_t *rects, int n_rects) +{ + int width = src->width; + int height = src->height; +#ifdef __ARM_NEON__ + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; +#endif + for (int i = 0; i < n_rects; i++) { + int left = (int)rects[i*4+0]; + int top = (int)rects[i*4+1]; + int right = (int)rects[i*4+2]; + int bottom = (int)rects[i*4+3]; + if (left < 0) left = 0; + if (top < 0) top = 0; + if (right > width) right = width; + if (bottom > height) bottom = height; + if (left >= right || top >= bottom) continue; + for (int row = top; row < bottom; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + int col = left; +#ifdef __ARM_NEON__ + // Advance scalar to the next multiple-of-8 boundary before the + // NEON loop. The only requirement for correct UV subsampling is + // that col is even; 8-alignment satisfies that and reduces the + // scalar pre-loop to at most 7 pixels (vs. 15 for 16-alignment). + int neon_start = (col + 7) & ~7; + for (; col < neon_start && col < right; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + for (; col + 15 < right; col += 16) + grdp_yuv420p_to_bgra_neon_16(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < right; col += 8) + grdp_yuv420p_to_bgra_neon_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); +#elif defined(__SSE2__) + // Advance to 8-aligned column so UV subsampling is always correct. + int sse_start = (col + 7) & ~7; + for (; col < sse_start && col < right; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + for (; col + 7 < right; col += 8) + grdp_yuv420p_to_bgra_sse2_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); +#endif + for (; col < right; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } + } +} + +// grdp_nv12_to_bgra converts a semi-planar NV12 frame (Y plane + interleaved +// UV plane) to packed BGRA using BT.601 coefficients. This bypasses swscale +// for the same reason as grdp_yuv420p_to_bgra: on ARM64 swscale's +// non-accelerated NV12→BGRA fallback ignores sws_setColorspaceDetails. +// VideoToolbox (macOS HW decoder) always outputs NV12. +// +// On ARM64 the inner loop is NEON-accelerated (8 pixels per iteration) to +// reduce per-frame CPU cost and decode-loop jitter. +#ifdef __ARM_NEON__ +// grdp_nv12_to_bgra_neon_8 processes 8 luma pixels (4 UV pairs) per call. +// All int32x4_t intermediates prevent overflow of e.g. 298*239 = 71 222. +static inline void grdp_nv12_to_bgra_neon_8( + const uint8_t *yrow, const uint8_t *uvrow, uint8_t *drow, + int col, int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, + int16_t yoff) +{ + // Load 8 luma bytes, convert to int16, subtract Y offset (16 or 0). + uint8x8_t y_u8 = vld1_u8(yrow + col); + int16x8_t c16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(y_u8)), + vdupq_n_s16(yoff)); + + // Load 8 UV bytes: [U0,V0,U1,V1,U2,V2,U3,V3]. + // vtrn_u8(a,a) transposes pairs: val[0]=[a[0],a[0],a[2],a[2],…] val[1]=[a[1],a[1],a[3],a[3],…]. + // Applied to interleaved NV12 this deinterleaves AND duplicates U and V in one instruction. + uint8x8_t uv_u8 = vld1_u8(uvrow + col); + uint8x8x2_t uv_dup = vtrn_u8(uv_u8, uv_u8); + uint8x8_t u8 = uv_dup.val[0]; // [U0,U0,U1,U1,U2,U2,U3,U3] + uint8x8_t v8 = uv_dup.val[1]; // [V0,V0,V1,V1,V2,V2,V3,V3] + + int16x8_t u16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u8)), vdupq_n_s16(128)); + int16x8_t v16 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v8)), vdupq_n_s16(128)); + + // Compute R/G/B with int32 to avoid overflow. Process 4+4 pixels. + int16x4_t c_lo = vget_low_s16(c16), u_lo = vget_low_s16(u16), v_lo = vget_low_s16(v16); + int16x4_t c_hi = vget_high_s16(c16), u_hi = vget_high_s16(u16), v_hi = vget_high_s16(v16); + + int32x4_t ky_lo = vmull_n_s16(c_lo, ky), ky_hi = vmull_n_s16(c_hi, ky); + + int32x4_t r_lo = vaddq_s32(vaddq_s32(ky_lo, vmull_n_s16(v_lo, kr)), vdupq_n_s32(128)); + int32x4_t g_lo = vaddq_s32(vsubq_s32(vsubq_s32(ky_lo, vmull_n_s16(u_lo, kgu)), vmull_n_s16(v_lo, kgv)), vdupq_n_s32(128)); + int32x4_t b_lo = vaddq_s32(vaddq_s32(ky_lo, vmull_n_s16(u_lo, kb)), vdupq_n_s32(128)); + + int32x4_t r_hi = vaddq_s32(vaddq_s32(ky_hi, vmull_n_s16(v_hi, kr)), vdupq_n_s32(128)); + int32x4_t g_hi = vaddq_s32(vsubq_s32(vsubq_s32(ky_hi, vmull_n_s16(u_hi, kgu)), vmull_n_s16(v_hi, kgv)), vdupq_n_s32(128)); + int32x4_t b_hi = vaddq_s32(vaddq_s32(ky_hi, vmull_n_s16(u_hi, kb)), vdupq_n_s32(128)); + + // Shift >>8, saturate int32→int16→uint8, then store interleaved BGRA. + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(r_lo,8)), vqmovn_s32(vshrq_n_s32(r_hi,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(g_lo,8)), vqmovn_s32(vshrq_n_s32(g_hi,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(b_lo,8)), vqmovn_s32(vshrq_n_s32(b_hi,8)))); + uint8x8x4_t bgra; + bgra.val[0] = b; + bgra.val[1] = g; + bgra.val[2] = r; + bgra.val[3] = vdup_n_u8(255); + vst4_u8(drow + col * 4, bgra); +} + +// grdp_nv12_to_bgra_neon_16 processes 16 luma pixels (8 UV pairs) per call. +// vld2_u8 deinterleaves U and V in one instruction; vzip_u8 then duplicates +// each chroma sample for the two luma pixels it serves, giving twice the +// throughput of grdp_nv12_to_bgra_neon_8 per loop iteration. +static inline void grdp_nv12_to_bgra_neon_16( + const uint8_t *yrow, const uint8_t *uvrow, uint8_t *drow, + int col, int16_t ky, int16_t kr, int16_t kgu, int16_t kgv, int16_t kb, + int16_t yoff) +{ + // Load 16 luma bytes, subtract Y offset. + uint8x16_t y_u8 = vld1q_u8(yrow + col); + int16x8_t c_lo = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(vget_low_u8(y_u8))), + vdupq_n_s16(yoff)); + int16x8_t c_hi = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(vget_high_u8(y_u8))), + vdupq_n_s16(yoff)); + + // Load 8 UV pairs (16 bytes) with automatic deinterleave. + // val[0]=[U0..U7], val[1]=[V0..V7]; each serves 2 luma pixels. + uint8x8x2_t uv = vld2_u8(uvrow + col); + uint8x8x2_t u_zip = vzip_u8(uv.val[0], uv.val[0]); // val[0]=[U0,U0,...,U3,U3] val[1]=[U4,U4,...,U7,U7] + uint8x8x2_t v_zip = vzip_u8(uv.val[1], uv.val[1]); + + // Pixels 0-7 (low half). + { + int16x8_t u8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u_zip.val[0])), vdupq_n_s16(128)); + int16x8_t v8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v_zip.val[0])), vdupq_n_s16(128)); + int16x4_t cl = vget_low_s16(c_lo), ul = vget_low_s16(u8), vl = vget_low_s16(v8); + int16x4_t ch = vget_high_s16(c_lo), uh = vget_high_s16(u8), vh = vget_high_s16(v8); + int32x4_t kyl = vmull_n_s16(cl, ky), kyh = vmull_n_s16(ch, ky); + int32x4_t rl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(vl, kr)), vdupq_n_s32(128)); + int32x4_t gl = vaddq_s32(vsubq_s32(vsubq_s32(kyl, vmull_n_s16(ul, kgu)), vmull_n_s16(vl, kgv)), vdupq_n_s32(128)); + int32x4_t bl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(ul, kb)), vdupq_n_s32(128)); + int32x4_t rh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(vh, kr)), vdupq_n_s32(128)); + int32x4_t gh = vaddq_s32(vsubq_s32(vsubq_s32(kyh, vmull_n_s16(uh, kgu)), vmull_n_s16(vh, kgv)), vdupq_n_s32(128)); + int32x4_t bh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(uh, kb)), vdupq_n_s32(128)); + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(rl,8)), vqmovn_s32(vshrq_n_s32(rh,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(gl,8)), vqmovn_s32(vshrq_n_s32(gh,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(bl,8)), vqmovn_s32(vshrq_n_s32(bh,8)))); + uint8x8x4_t bgra; bgra.val[0]=b; bgra.val[1]=g; bgra.val[2]=r; bgra.val[3]=vdup_n_u8(255); + vst4_u8(drow + col * 4, bgra); + } + // Pixels 8-15 (high half). + { + int16x8_t u8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(u_zip.val[1])), vdupq_n_s16(128)); + int16x8_t v8 = vsubq_s16(vreinterpretq_s16_u16(vmovl_u8(v_zip.val[1])), vdupq_n_s16(128)); + int16x4_t cl = vget_low_s16(c_hi), ul = vget_low_s16(u8), vl = vget_low_s16(v8); + int16x4_t ch = vget_high_s16(c_hi), uh = vget_high_s16(u8), vh = vget_high_s16(v8); + int32x4_t kyl = vmull_n_s16(cl, ky), kyh = vmull_n_s16(ch, ky); + int32x4_t rl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(vl, kr)), vdupq_n_s32(128)); + int32x4_t gl = vaddq_s32(vsubq_s32(vsubq_s32(kyl, vmull_n_s16(ul, kgu)), vmull_n_s16(vl, kgv)), vdupq_n_s32(128)); + int32x4_t bl = vaddq_s32(vaddq_s32(kyl, vmull_n_s16(ul, kb)), vdupq_n_s32(128)); + int32x4_t rh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(vh, kr)), vdupq_n_s32(128)); + int32x4_t gh = vaddq_s32(vsubq_s32(vsubq_s32(kyh, vmull_n_s16(uh, kgu)), vmull_n_s16(vh, kgv)), vdupq_n_s32(128)); + int32x4_t bh = vaddq_s32(vaddq_s32(kyh, vmull_n_s16(uh, kb)), vdupq_n_s32(128)); + uint8x8_t r = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(rl,8)), vqmovn_s32(vshrq_n_s32(rh,8)))); + uint8x8_t g = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(gl,8)), vqmovn_s32(vshrq_n_s32(gh,8)))); + uint8x8_t b = vqmovun_s16(vcombine_s16(vqmovn_s32(vshrq_n_s32(bl,8)), vqmovn_s32(vshrq_n_s32(bh,8)))); + uint8x8x4_t bgra; bgra.val[0]=b; bgra.val[1]=g; bgra.val[2]=r; bgra.val[3]=vdup_n_u8(255); + vst4_u8(drow + (col + 8) * 4, bgra); + } +} +#endif // __ARM_NEON__ + +static void grdp_nv12_to_bgra( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range) +{ + int width = src->width; + int height = src->height; +#ifdef __ARM_NEON__ + // NEON fast path: 8 pixels per inner iteration on ARM64. + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 15 < width; col += 16) + grdp_nv12_to_bgra_neon_16(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < width; col += 8) + grdp_nv12_to_bgra_neon_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + // Scalar tail for widths not a multiple of 8. + for (; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 7 < width; col += 8) + grdp_nv12_to_bgra_sse2_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#else + for (int row = 0; row < height; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + for (int col = 0; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#endif +} + +// grdp_nv12_to_bgra_regions is the region-aware variant of grdp_nv12_to_bgra. +// Only pixels within the n_rects dirty rectangles (flat [left,top,right,bottom] +// uint16 tuples) are written; all other pixels in dst are left untouched. +static void grdp_nv12_to_bgra_regions( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range, + const uint16_t *rects, int n_rects) +{ + int width = src->width; + int height = src->height; +#ifdef __ARM_NEON__ + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; +#endif + for (int i = 0; i < n_rects; i++) { + int left = (int)rects[i*4+0]; + int top = (int)rects[i*4+1]; + int right = (int)rects[i*4+2]; + int bottom = (int)rects[i*4+3]; + if (left < 0) left = 0; + if (top < 0) top = 0; + if (right > width) right = width; + if (bottom > height) bottom = height; + if (left >= right || top >= bottom) continue; + for (int row = top; row < bottom; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + int col = left; +#ifdef __ARM_NEON__ + int neon_start = (col + 7) & ~7; + for (; col < neon_start && col < right; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + for (; col + 15 < right; col += 16) + grdp_nv12_to_bgra_neon_16(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < right; col += 8) + grdp_nv12_to_bgra_neon_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); +#elif defined(__SSE2__) + int sse_start = (col + 7) & ~7; + for (; col < sse_start && col < right; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + for (; col + 7 < right; col += 8) + grdp_nv12_to_bgra_sse2_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); +#endif + for (; col < right; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } + } +} + +// grdp_sample_yuv samples the centre pixel of a planar YUV frame for +// diagnostics. Returns raw byte values (not offset-adjusted). +static void grdp_sample_yuv(const AVFrame *f, + uint8_t *y_out, uint8_t *u_out, uint8_t *v_out) +{ + int cx = f->width / 2; + int cy = f->height / 2; + *y_out = f->data[0][ cy * f->linesize[0] + cx ]; + *u_out = f->data[1][(cy / 2) * f->linesize[1] + (cx / 2)]; + *v_out = f->data[2][(cy / 2) * f->linesize[2] + (cx / 2)]; +} + +// grdp_sample_nv12 samples the centre pixel of a semi-planar NV12 frame +// (Y plane + interleaved UV plane) for diagnostics. +static void grdp_sample_nv12(const AVFrame *f, + uint8_t *y_out, uint8_t *u_out, uint8_t *v_out) +{ + int cx = f->width / 2; + int cy = f->height / 2; + *y_out = f->data[0][cy * f->linesize[0] + cx]; + int uvx = (cx / 2) * 2; // NV12: interleaved, U at even index, V at odd + int uvy = cy / 2; + *u_out = f->data[1][uvy * f->linesize[1] + uvx]; + *v_out = f->data[1][uvy * f->linesize[1] + uvx + 1]; +} + +// grdp_sample_nv12_at samples a specific (x,y) pixel from an NV12 frame. +static void grdp_sample_nv12_at(const AVFrame *f, int x, int y, + uint8_t *y_out, uint8_t *u_out, uint8_t *v_out) +{ + if (x < 0 || x >= f->width || y < 0 || y >= f->height) { + *y_out = *u_out = *v_out = 0; + return; + } + *y_out = f->data[0][y * f->linesize[0] + x]; + int uvx = (x >> 1) * 2; + int uvy = y >> 1; + *u_out = f->data[1][uvy * f->linesize[1] + uvx]; + *v_out = f->data[1][uvy * f->linesize[1] + uvx + 1]; +} + +// grdp_is_warmup_nv12 checks whether an NV12 frame looks like an uninitialised +// VideoToolbox IOSurface where the Y plane is zeroed but UV is at neutral 128. +// It samples Y at a 3×3 grid spread across the frame (at 25%, 50%, 75% of +// width and height). Returns 1 if every sampled luma value is <= threshold. +// A threshold of 4 catches Y=0 warm-up frames while tolerating H.264 rounding +// (luma=1 is possible). Real video content — even a mostly-black desktop — is +// statistically unlikely to have all 9 spread-out luma samples at or below 4. +static int grdp_is_warmup_nv12(const AVFrame *f, int threshold) { + const uint8_t *yp = f->data[0]; + int stride = f->linesize[0]; + int w = f->width, h = f->height; + if (!yp || w <= 0 || h <= 0) return 0; + int xs[3] = { w / 4, w / 2, 3 * w / 4 }; + int ys[3] = { h / 4, h / 2, 3 * h / 4 }; + for (int i = 0; i < 3; i++) + for (int j = 0; j < 3; j++) + if (yp[ys[i] * stride + xs[j]] > (uint8_t)threshold) + return 0; + return 1; +} + +// grdp_debug_sample_y_grid fills out[0..8] with the same 3x3 luma samples +// grdp_is_warmup_nv12 checks (25%/50%/75% of width and height), for +// diagnostic logging only — it does not affect the warm-up decision itself. +static void grdp_debug_sample_y_grid(const AVFrame *f, uint8_t *out) { + const uint8_t *yp = f->data[0]; + int stride = f->linesize[0]; + int w = f->width, h = f->height; + if (!yp || w <= 0 || h <= 0) { + for (int k = 0; k < 9; k++) out[k] = 0; + return; + } + int xs[3] = { w / 4, w / 2, 3 * w / 4 }; + int ys[3] = { h / 4, h / 2, 3 * h / 4 }; + for (int i = 0; i < 3; i++) + for (int j = 0; j < 3; j++) + out[i * 3 + j] = yp[ys[i] * stride + xs[j]]; +} + +// grdp_is_low_chroma_nv12 returns 1 if at least minLow of the sampled UV pairs +// in an NV12 frame are abnormally low (both U and V below threshold). Valid +// YUV content always has chroma centered near 128; corruption from stale IDR +// priming or uninitialised buffers collapses chroma to ~0, which renders as a +// green-monochrome frame. A 3x3 grid spread across the frame catches corruption +// that only affects the centre pixel. +static int grdp_is_low_chroma_nv12(const AVFrame *f, int threshold, int minLow) { + const uint8_t *uvp = f->data[1]; + int stride = f->linesize[1]; + int w = f->width, h = f->height; + if (!uvp || w <= 0 || h <= 0) return 0; + int ph = (h + 1) / 2; + int pw = (w + 1) / 2; + if (pw <= 0 || ph <= 0) return 0; + int xs[3] = { pw / 4, pw / 2, 3 * pw / 4 }; + int ys[3] = { ph / 4, ph / 2, 3 * ph / 4 }; + int low = 0; + for (int i = 0; i < 3; i++) { + for (int j = 0; j < 3; j++) { + int x = xs[j] * 2; // NV12: interleaved U at even byte, V at odd byte + int y = ys[i]; + int u = (int)uvp[y * stride + x]; + int v = (int)uvp[y * stride + x + 1]; + if (u < threshold && v < threshold) low++; + } + } + return low >= minLow ? 1 : 0; +} + +// grdp_is_low_chroma_yuv420p is the planar YUV420P equivalent of +// grdp_is_low_chroma_nv12. +static int grdp_is_low_chroma_yuv420p(const AVFrame *f, int threshold, int minLow) { + const uint8_t *up = f->data[1]; + const uint8_t *vp = f->data[2]; + int uStride = f->linesize[1]; + int vStride = f->linesize[2]; + int w = f->width, h = f->height; + if (!up || !vp || w <= 0 || h <= 0) return 0; + int ph = (h + 1) / 2; + int pw = (w + 1) / 2; + if (pw <= 0 || ph <= 0) return 0; + int xs[3] = { pw / 4, pw / 2, 3 * pw / 4 }; + int ys[3] = { ph / 4, ph / 2, 3 * ph / 4 }; + int low = 0; + for (int i = 0; i < 3; i++) { + for (int j = 0; j < 3; j++) { + int x = xs[j]; + int y = ys[i]; + int u = (int)up[y * uStride + x]; + int v = (int)vp[y * vStride + x]; + if (u < threshold && v < threshold) low++; + } + } + return low >= minLow ? 1 : 0; +} + +static void grdp_sws_set_src_range(struct SwsContext *sws, int full_range) { + const int *inv_table, *table; + int src_range, dst_range, brightness, contrast, saturation; + if (sws_getColorspaceDetails(sws, + (int **)&inv_table, &src_range, + (int **)&table, &dst_range, + &brightness, &contrast, &saturation) >= 0) { + sws_setColorspaceDetails(sws, + inv_table, full_range, + table, dst_range, + brightness, contrast, saturation); + } +} + +// grdp_copy_yuv420p_to_i420 copies an AVFrame in YUV420P or YUVJ420P format +// to tightly-packed I420 planes (stride = width for Y, stride = (width+1)/2 for U/V). +// ydst, udst, vdst must be pre-allocated by the caller. +// When the source frame is tight-packed (linesize == stride), a single bulk +// memcpy is used per plane instead of per-row copies. +static void grdp_copy_yuv420p_to_i420( + const AVFrame *f, + uint8_t *ydst, uint8_t *udst, uint8_t *vdst, + int w, int h) +{ + int pw = (w + 1) / 2; + int ph = (h + 1) / 2; + if (f->linesize[0] == w) + memcpy(ydst, f->data[0], (size_t)w * h); + else + for (int y = 0; y < h; y++) + memcpy(ydst + y * w, f->data[0] + y * f->linesize[0], w); + if (f->linesize[1] == pw) + memcpy(udst, f->data[1], (size_t)pw * ph); + else + for (int y = 0; y < ph; y++) + memcpy(udst + y * pw, f->data[1] + y * f->linesize[1], pw); + if (f->linesize[2] == pw) + memcpy(vdst, f->data[2], (size_t)pw * ph); + else + for (int y = 0; y < ph; y++) + memcpy(vdst + y * pw, f->data[2] + y * f->linesize[2], pw); +} + +// grdp_copy_nv12_to_i420 copies an AVFrame in NV12 format (Y plane + interleaved UV) +// to tightly-packed I420 planes. +// The Y plane is bulk-copied when tight-packed. +// On ARM64 the UV deinterleave loop uses NEON vld2q_u8 to process 16 chroma +// pairs per iteration, roughly halving the cost of the chroma plane copy. +static void grdp_copy_nv12_to_i420( + const AVFrame *f, + uint8_t *ydst, uint8_t *udst, uint8_t *vdst, + int w, int h) +{ + int pw = (w + 1) / 2; + int ph = (h + 1) / 2; + if (f->linesize[0] == w) + memcpy(ydst, f->data[0], (size_t)w * h); + else + for (int y = 0; y < h; y++) + memcpy(ydst + y * w, f->data[0] + y * f->linesize[0], w); + for (int y = 0; y < ph; y++) { + const uint8_t *row = f->data[1] + y * f->linesize[1]; + uint8_t *ud = udst + y * pw; + uint8_t *vd = vdst + y * pw; + int x = 0; +#ifdef __ARM_NEON__ + for (; x + 15 < pw; x += 16) { + uint8x16x2_t uv = vld2q_u8(row + x * 2); + vst1q_u8(ud + x, uv.val[0]); + vst1q_u8(vd + x, uv.val[1]); + } +#elif defined(__SSE2__) + // SSE2: deinterleave 8 UV pairs (16 bytes) per iteration. + // _mm_and_si128 extracts even bytes (U) and _mm_srli_epi16 shifts odd bytes (V). + // _mm_packus_epi16 packs 8 × 16-bit values to 8 bytes in the low half. + for (; x + 7 < pw; x += 8) { + __m128i uv128 = _mm_loadu_si128((const __m128i *)(row + x * 2)); + __m128i u_vec = _mm_and_si128(uv128, _mm_set1_epi16(0x00FF)); + __m128i v_vec = _mm_srli_epi16(uv128, 8); + _mm_storel_epi64((__m128i *)(ud + x), _mm_packus_epi16(u_vec, u_vec)); + _mm_storel_epi64((__m128i *)(vd + x), _mm_packus_epi16(v_vec, v_vec)); + } +#endif + for (; x < pw; x++) { + ud[x] = row[x * 2]; + vd[x] = row[x * 2 + 1]; + } + } +} + +// grdp_copy_nv12 copies an AVFrame in NV12 format to tightly-packed NV12 +// planes. Unlike grdp_copy_nv12_to_i420, this keeps the interleaved UV plane +// intact so SDL2 can upload it via SDL_UpdateNVTexture without CPU-side +// deinterleaving. +// When the source frame is tight-packed (linesize == stride), a single bulk +// memcpy is used per plane instead of per-row copies. +static void grdp_copy_nv12( + const AVFrame *f, + uint8_t *ydst, uint8_t *uvdst, + int w, int h) +{ + int uv_bytes = ((w + 1) / 2) * 2; + int ph = (h + 1) / 2; + if (f->linesize[0] == w) + memcpy(ydst, f->data[0], (size_t)w * h); + else + for (int y = 0; y < h; y++) + memcpy(ydst + y * w, f->data[0] + y * f->linesize[0], w); + if (f->linesize[1] == uv_bytes) + memcpy(uvdst, f->data[1], (size_t)uv_bytes * ph); + else + for (int y = 0; y < ph; y++) + memcpy(uvdst + y * uv_bytes, f->data[1] + y * f->linesize[1], uv_bytes); +} + +// grdp_nv12_to_bgra_rows is the row-range variant of grdp_nv12_to_bgra. +// Only rows [start_row, end_row) are written; the rest of dst is untouched. +// This allows the caller to parallelise the conversion across goroutines. +static void grdp_nv12_to_bgra_rows( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range, + int start_row, int end_row) +{ + int width = src->width; + if (start_row < 0) start_row = 0; + if (end_row > src->height) end_row = src->height; +#ifdef __ARM_NEON__ + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 15 < width; col += 16) + grdp_nv12_to_bgra_neon_16(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < width; col += 8) + grdp_nv12_to_bgra_neon_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 7 < width; col += 8) + grdp_nv12_to_bgra_sse2_8(yrow, uvrow, drow, col, ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#else + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *uvrow = src->data[1] + (row >> 1) * src->linesize[1]; + uint8_t *drow = dst + row * dst_stride; + for (int col = 0; col < width; col++) { + int u = (int)uvrow[(col >> 1) * 2 ] - 128; + int v = (int)uvrow[(col >> 1) * 2 + 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#endif +} + +// grdp_yuv420p_to_bgra_rows is the row-range variant of grdp_yuv420p_to_bgra. +static void grdp_yuv420p_to_bgra_rows( + const AVFrame *src, uint8_t *dst, int dst_stride, int full_range, + int start_row, int end_row) +{ + int width = src->width; + if (start_row < 0) start_row = 0; + if (end_row > src->height) end_row = src->height; +#ifdef __ARM_NEON__ + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 15 < width; col += 16) + grdp_yuv420p_to_bgra_neon_16(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col + 7 < width; col += 8) + grdp_yuv420p_to_bgra_neon_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#elif defined(__SSE2__) + int16_t ky = full_range ? 256 : 298; + int16_t kr = full_range ? 359 : 409; + int16_t kgu = full_range ? 88 : 100; + int16_t kgv = full_range ? 183 : 208; + int16_t kb = full_range ? 454 : 516; + int16_t yoff = full_range ? 0 : 16; + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + int col = 0; + for (; col + 7 < width; col += 8) + grdp_yuv420p_to_bgra_sse2_8(yrow, urow, vrow, drow, col, + ky, kr, kgu, kgv, kb, yoff); + for (; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#else + for (int row = start_row; row < end_row; row++) { + const uint8_t *yrow = src->data[0] + row * src->linesize[0]; + const uint8_t *urow = src->data[1] + (row >> 1) * src->linesize[1]; + const uint8_t *vrow = src->data[2] + (row >> 1) * src->linesize[2]; + uint8_t *drow = dst + row * dst_stride; + for (int col = 0; col < width; col++) { + int u = (int)urow[col >> 1] - 128; + int v = (int)vrow[col >> 1] - 128; + grdp_bt601_pixel((int)yrow[col], u, v, full_range, drow + col*4); + } + } +#endif +} + +// grdp_find_v4l2m2m returns the h264_v4l2m2m decoder if FFmpeg was compiled +// with V4L2 M2M support (common on Linux/Raspberry Pi). Returns NULL otherwise. +static const AVCodec *grdp_find_v4l2m2m(void) { + return avcodec_find_decoder_by_name("h264_v4l2m2m"); +} +*/ +import "C" + +import ( + "fmt" + "log/slog" + "runtime" + "sync" + "sync/atomic" + "time" + "unsafe" + + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpgfx" +) + +// useSwscale controls whether YUV420P/YUVJ420P and NV12 frames are converted +// to BGRA via swscale (SIMD-accelerated on x86_64) or via hand-written C loops. +// On ARM64, swscale's non-accelerated paths ignore sws_setColorspaceDetails, +// producing a strong green cast; the hand-written BT.601 loops are used instead. +// On x86_64, swscale is both correct and significantly faster (SSSE3/AVX2). +var useSwscale = runtime.GOARCH != "arm64" + +// convertWorkers is the number of goroutines used to parallelise the YUV→BGRA +// conversion. Capped at 4 so we don't thrash the cache with too many threads +// writing into the same output buffer. +var convertWorkers = min(runtime.GOMAXPROCS(0), 4) + +// convertParallelMinH is the minimum frame height at which row-parallel +// conversion is enabled. For small frames the goroutine overhead outweighs +// the gain from parallelism. +const convertParallelMinH = 480 + +// rowConvertJob is the work item dispatched to convertWorkerPool goroutines. +type rowConvertJob struct { + fn func(s, e C.int) + s, e C.int +} + +// convertWorkerPool maintains N persistent goroutines for row-parallel +// YUV→BGRA conversion, eliminating the per-frame goroutine allocation cost. +// At 30fps with 4 workers the old code created ~120 goroutines/second; this +// pool amortises that overhead down to zero after the first frame. +type convertWorkerPool struct { + jobs chan rowConvertJob + done chan struct{} + n int +} + +func newConvertWorkerPool(n int) *convertWorkerPool { + p := &convertWorkerPool{ + jobs: make(chan rowConvertJob, n), + done: make(chan struct{}, n), + n: n, + } + for range n { + go func() { + for job := range p.jobs { + job.fn(job.s, job.e) + p.done <- struct{}{} + } + }() + } + return p +} + +// dispatch partitions [0, h) into p.n equal bands and executes fn +// concurrently across the pool's persistent goroutines. +func (p *convertWorkerPool) dispatch(h int, fn func(s, e C.int)) { + step := (h + p.n - 1) / p.n + count := 0 + for i := 0; i < p.n; i++ { + s := i * step + e := s + step + if e > h { + e = h + } + if s >= h { + break + } + p.jobs <- rowConvertJob{fn: fn, s: C.int(s), e: C.int(e)} + count++ + } + for range count { + <-p.done + } +} + +var ( + globalConvertPool *convertWorkerPool + globalConvertPoolOnce sync.Once +) + +// parallelConvertRows partitions [0, h) into convertWorkers equal bands and +// calls fn(startRow, endRow) concurrently via a persistent worker pool. For +// small frames (h < convertParallelMinH) or when convertWorkers <= 1 the +// function is called serially to avoid scheduling overhead. +func parallelConvertRows(h int, fn func(s, e C.int)) { + n := convertWorkers + if n <= 1 || h < convertParallelMinH { + fn(0, C.int(h)) + return + } + globalConvertPoolOnce.Do(func() { + globalConvertPool = newConvertWorkerPool(n) + }) + globalConvertPool.dispatch(h, fn) +} + +// avLogOnce ensures grdp_suppress_av_log is called only once per process. +var avLogOnce sync.Once + +// avcFreezeThreshold is the duration of no decoded output from the HW decoder +// lowChromaThreshold is the UV value below which a chroma sample is considered +// abnormally low. Valid YUV content centres chroma near 128; corruption from +// stale IDR priming or uninitialised buffers collapses chroma toward 0. The +// observed SW-fallback corruption produces U≈60, V≈42, so 72 catches those +// green-monochrome frames while staying well below the U≈101–122, V≈88–122 +// range seen in healthy desktop content. +const lowChromaThreshold = 72 + +// after which it is marked broken. The application-level watchdog then +// reconnects the RDP session. VideoToolbox (macOS) can stall for 2-3 seconds +// while processing a new IDR/SPS frame; 6 seconds gives it enough headroom to +// recover naturally before we declare it broken. FreeRDP takes a similar +// passive approach: it drops failed frames without hard resets or IDR requests +// and waits for the server to resume naturally. +// This threshold applies to the initial-stall case (hwReady=false). +const avcFreezeThreshold = 6 * time.Second + +// avcHWReadyFreezeThreshold is the point at which a stalled HW decoder stops +// accepting new packets. VideoToolbox can legitimately pause for several +// seconds when flushing its internal reference pipeline at an IDR/GOP +// boundary; 5 s is chosen because the CGo call (avcodec_send_packet) itself +// permanently blocks after ~5.75 s of stall on macOS VideoToolbox. The +// pre-flight guard in Decode() bails out *before* the CGo call to prevent the +// decodeLoop goroutine from hanging inside CGo. +// +// Crossing this threshold does NOT immediately mark the decoder broken. +// Instead Decode() enters a recovery-probe window (avcHWRecoveryWindow) and +// keeps probing avcodec_receive_frame without sending new packets. If +// VideoToolbox produces a frame during that window the stall clock is reset +// and normal decoding resumes; only if the window is exhausted is the decoder +// marked broken. +const avcHWReadyFreezeThreshold = 5 * time.Second + +// avcHWEarlyFreezeThreshold is a shorter stall threshold applied during the +// first avcHWEarlyFrameLimit packets sent to the HW decoder after each +// decoder initialisation or flush. VideoToolbox exhibits a characteristic +// stall pattern at RDP session start (and after a forced flush): it processes +// a small burst of frames, then freezes for 8+ seconds without recovering. +// The normal 7 s threshold is designed for mid-session IDR stalls that +// self-resolve in 2-3 s; in the early phase a genuine VT stall is +// distinguishable because it persists well beyond 5 s. Using a shorter +// threshold here reduces the visible freeze by ~2 s while avoiding false +// positives from transient 3-4 s null-frame bursts at IDR/GOP boundaries. +const avcHWEarlyFreezeThreshold = 5 * time.Second + +// avcHWEarlyFrameLimit is the number of packets sent to the HW decoder +// (hwSentCount) below which avcHWEarlyFreezeThreshold is used instead of +// avcHWReadyFreezeThreshold. hwSentCount resets to zero on every +// avcodec_send_packet failure (decoder flush), so this threshold also covers +// the early window after an in-session flush. 50 packets at 30 fps ≈ 1.7 s, +// comfortably covering the unstable session-start window without interfering +// with normal mid-session IDR stalls. +const avcHWEarlyFrameLimit = 50 + +// avcHWRecoveryWindow is how long Decode() probes for pending output after +// a stall is detected (either from avcHWReadyFreezeThreshold or from the +// early null-frame count detector). After this window, if VT has not +// produced a real frame, the decoder is declared broken and a soft reset +// is triggered. 300 ms is sufficient: any frame VT had buffered will have +// surfaced well within 300 ms, and YouTube / gnome-remote-desktop delivers a +// fresh IDR within ~2 s of the soft reset anyway. +const avcHWRecoveryWindow = 300 * time.Millisecond + +// avcHWNullFrameStallLimit is the number of consecutive null (blank) frames +// the HW decoder may produce during the early session window (hwSentCount < +// avcHWEarlyFrameLimit) before triggering a stall-probe. VideoToolbox +// legitimately outputs a handful of null frames at IDR boundaries during +// session start (observed: 5–25 frames / up to ~1 s at 30 fps). +const avcHWNullFrameStallLimit = 25 + +// avcHWEarlyStallMinElapsed is the minimum wall-clock time since the first +// packet was sent (hwFirstSendTime) before the early-window null-frame stall +// probe is allowed to fire. VideoToolbox needs roughly 1 s to initialise +// its pipeline regardless of how fast packets arrive, so at high packet rates +// (e.g. a burst at session start) the 25-frame count threshold is hit in +// milliseconds — far too quickly to distinguish a genuine stall from normal +// initialisation. By requiring at least 2 s of elapsed time we avoid +// false-positive probes while still detecting real stalls well before the +// 7-second safety valve. +const avcHWEarlyStallMinElapsed = 2 * time.Second + +// avcHWMidSessionNullFrameLimit is the null-frame count threshold used for +// mid-session stall detection (hwSentCount >= avcHWEarlyFrameLimit). +// Normal GOP / mid-session IDR boundaries can produce up to ~25 null frames +// (~1 s at 30 fps); a genuine VideoToolbox stall produces hundreds. 15 +// frames ≈ 0.5 s at 30 fps — this is more aggressive than the previous 20 +// and may occasionally trigger a SW fallback at a noisy GOP boundary, but +// it recovers from genuine VT stalls noticeably faster. Observed logs show +// zero mid-session null frames outside of the fatal VT stall, so the trade-off +// is acceptable for the target use case. +const avcHWMidSessionNullFrameLimit = 15 + +// avcHWConsecDroppedFrameLimit is the number of consecutive HW-decoded +// frames that may be classified as warm-up/zero-fill/low-chroma junk (see +// the hwNeedsZeroCheck and low-chroma guards in convertFrame) before the +// decoder is declared broken, even though avcodec_receive_frame keeps +// returning frames (so hwConsecNullFrames / lastSuccessTime never reflect a +// stall). Under normal operation hwNeedsZeroCheck disarms itself after the +// first non-dropped frame, so this streak should only ever be a handful of +// frames long; a genuine VideoToolbox malfunction can instead keep emitting +// black/zero content indefinitely, leaving the on-screen video frozen on a +// black frame while the existing stall detectors — which only watch for the +// *absence* of frames — stay quiet. 150 frames is ~5 s at 30 fps. +const avcHWConsecDroppedFrameLimit = 150 + +// avcHWDroppedFrameStallThreshold is the elapsed-time companion to +// avcHWConsecDroppedFrameLimit: even at a low or irregular frame rate, a +// dropped-frame streak lasting this long is not a normal warm-up burst. +const avcHWDroppedFrameStallThreshold = 5 * time.Second + +// avcHWConsecDroppedFrameMinCount is the minimum streak length required +// before avcHWDroppedFrameStallThreshold is honoured, so a single slow +// frame arriving 5 s after the previous one cannot trip the time-based leg +// on its own. +const avcHWConsecDroppedFrameMinCount = 10 + +// keyframeWaitLimit is the maximum number of non-IDR packets we drop while +// waiting for a keyframe after a decoder reset or flush. gnome-remote-desktop +// and similar servers send an IDR approximately every 15-25 seconds; using 900 +// frames (~30s at 30 fps) ensures we catch the next natural IDR even under +// variable server GOP intervals. After this limit the SW decoder attempts +// error-concealment decode; the HW decoder marks itself broken instead (HW +// codecs like VideoToolbox cannot recover without a proper IDR). +const keyframeWaitLimit = 900 + +// keyframeWaitTimeout is the maximum wall-clock time an HW decoder waits for +// an IDR after entering needsKeyFrame=true. ForceRefresh is sent every 2 s, +// so the server should respond within a few seconds. If no IDR arrives within +// this window the HW decoder marks itself broken so the soft-reset / reconnect +// chain escalates quickly rather than waiting the full keyframeWaitLimit +// (~30 s) of dropped packets. The HW path cannot do error-concealment, so a +// longer window gives the server more chances to deliver a natural IDR. +const keyframeWaitTimeout = 15 * time.Second + +// keyframeWaitTimeoutSW is the keyframe-wait wall-clock limit for detached SW +// decoders (aux decoder h264dec2, no watchdog channel). These decoders do not +// have an external timer to terminate the wait, so we attempt error-concealment +// sooner and let the caller tear down and recreate the decoder on the next IDR. +const keyframeWaitTimeoutSW = 5 * time.Second + +// keyframeWaitTimeoutSWFallback is the keyframe-wait timer for the main SW +// fallback decoder (created after a VideoToolbox stall). Set to 3 s to give +// Windows Server enough time to respond to the ForceRefresh (SuppressOutput +// toggle) with a fresh IDR before escalating to a full reconnect. Some +// Windows Server versions respond within 1–2 s; others never respond in +// AVC444 mode, in which case the reconnect happens after 3 s regardless. +// The 3 s window avoids spurious reconnects when the server responds slowly. +const keyframeWaitTimeoutSWFallback = 3 * time.Second + +// profileWindow is the number of HW frames over which Decode aggregates +// timing measurements before logging an INFO summary. At 30 fps this is +// roughly one log line every ~10 s. +const profileWindow = 300 + +type ffmpegDecoder struct { + codecCtx *C.AVCodecContext + packet *C.AVPacket + frame *C.AVFrame + swFrame *C.AVFrame + mapFrame *C.AVFrame // reusable frame for av_hwframe_map zero-copy transfers + hwMapSupported int8 // 0=unknown, 1=supported, -1=unsupported + swsCtx *C.struct_SwsContext + useHW bool + hwPixFmt C.enum_AVPixelFormat + lastW C.int + lastH C.int + lastFmt C.enum_AVPixelFormat + lastFullRange C.int // tracks fullRange used when swsCtx was last configured + lastSuccessTime time.Time // wall-clock time of the last successfully decoded frame + lastSendTime time.Time // wall-clock time of the last avcodec_send_packet call + lastReceiveTime time.Time // wall-clock time of the last Decode() call (updated on every call) + hwFirstSendTime time.Time // wall-clock time of the first packet sent to the HW decoder + needsKeyFrame bool // drop packets until an IDR/SPS is received + keyframeWaitCount int // P-frames dropped so far while needsKeyFrame=true + keyframeWaitStart time.Time // wall-clock time of the first dropped P-frame while waiting for IDR + hwReady bool // HW decoder has produced at least one frame + hwSentCount int // packets sent to HW decoder (for diagnostics) + swFrameCount int // frames decoded by SW decoder (for diagnostics) + hwFrameCount int // frames decoded by HW decoder (for diagnostics) + broken bool // decoder is unrecoverable; stop producing frames so the app reconnects + brokenReason rdpgfx.H264BrokenReason + timerBroken atomic.Bool // set by background timers when probe/IDR timeouts expire + timerBrokenReason atomic.Int32 + proceededWithoutKeyframe bool // "proceed without keyframe" path was taken; AVERROR here means broken + stallProbeStart time.Time // wall-clock time we entered the stall recovery-probe window + stallTimer *time.Timer // fires after avcHWRecoveryWindow to mark broken independently of frame rate + hwConsecNullFrames int // consecutive HW null frames since last real frame; for early stall detection + hwConsecDroppedFrames int // consecutive HW frames classified as warm-up/zero-fill/low-chroma and dropped + hwFirstDroppedTime time.Time // wall-clock time the current dropped-frame streak began + kfWaitTimer *time.Timer // fires after kfWaitTimeoutVal to mark broken independently of frame rate + kfWaitTimeoutVal time.Duration // per-decoder IDR wait limit (varies by decoder type) + watchdogCh chan<- struct{} // signals GfxHandler.decodeLoop to call maybeNotifyDecoderBroken + + // Profiling: aggregated timing stats over the last profileWindow frames + // for the HW path. Helps determine whether convertFrame + // (av_hwframe_transfer_data + colour conversion) is the bottleneck that + // causes VideoToolbox to stall by holding GPU frames too long. + profFrames int + profSendNs int64 // total ns in avcodec_send_packet + profRecvNs int64 // total ns in avcodec_receive_frame loop (excluding convert) + profConvertNs int64 // total ns in convertFrame (transfer + colour conversion) + profTransferNs int64 // total ns in av_hwframe_transfer_data only + profMaxConvNs int64 // worst-case convertFrame duration in window + profMaxSendNs int64 // worst-case avcodec_send_packet duration in window + profMaxRecvNs int64 // worst-case avcodec_receive_frame duration in window + + // outRing holds two recyclable BGRA destination buffers. convertFrame + // rotates between them so each Decode() avoids allocating a fresh + // width*height*4 buffer (≈8MB at 1920×1080 → ≈240MB/s of GC garbage at + // 30fps). Two slots is sufficient because emitBitmap is called + // synchronously from the rdpgfx PDU loop and always finishes (the + // caller has copied the data into its backing image) before the next + // Decode runs. outRingIdx selects the slot to use *next*. + outRing [2][]byte + outRingIdx int + + // outI420Ring holds two recyclable I420 frame slots for GPU-accelerated + // rendering via SDL2 IYUV textures. Same ring/lifecycle pattern as outRing. + // outI420Enabled gates I420 extraction (set by DecodeWithI420); outNV12Enabled + // also triggers I420 extraction for non-NV12 sources (e.g. software YUV420P) + // so the caller can use LastI420() to update the AVC444 Y cache. lastI420 + // is the result from the most recent convertFrame call. + outI420Ring [2]rdpgfx.H264FrameI420 + outI420RingIdx int + outI420Enabled bool + lastI420 *rdpgfx.H264FrameI420 + + // outNV12Ring holds native NV12 frames for SDL2 NV12 texture upload. + // This path is especially useful for VideoToolbox, whose transferred + // software frames are usually NV12. + outNV12Ring [2]rdpgfx.H264FrameNV12 + outNV12RingIdx int + outNV12Enabled bool + lastNV12 *rdpgfx.H264FrameNV12 + + // regionHint carries dirty-rect hints for region-aware YUV→BGRA conversion. + // setRegionHint populates these fields; Decode() captures them into local + // variables at entry (clearing nRegionHints) so stale hints can never + // carry over to a subsequent unrelated frame. + regionHint []C.uint16_t // flat [left,top,right,bottom,...] per rect + nRegionHints C.int // number of valid rects in regionHint + + // hwNeedsZeroCheck is set to true on decoder creation and after each + // avcodec_flush_buffers call. When set, convertFrame checks the first + // NV12 output frame for a zero-filled chroma plane (U=0, V=0). + // + // VideoToolbox sometimes returns a zero-initialised IOSurface for the + // first decoded frame after init or a pipeline flush. The BT.601 + // limited-range conversion of (Y=0, U=0, V=0) produces BGRA(0,135,0,255) + // — a full-screen dark-green frame that manifests as a brief "green + // curtain" in the UI. Valid NV12 chroma always centres on 128 + // (limited-range [16,240], full-range centred on 128), so U=0 and V=0 + // occurring simultaneously at the centre pixel unambiguously signals an + // uninitialised buffer rather than real video content. + hwNeedsZeroCheck bool + + // swNeedsZeroCheck mirrors hwNeedsZeroCheck for the software (YUV420P) + // decode path. The FFmpeg SW decoder similarly outputs all-zero frames + // on the first packet after creation or avcodec_flush_buffers, especially + // when primed with a stale cached IDR. Cleared after the first valid + // (non-zero-fill) YUV420P frame is seen. + swNeedsZeroCheck bool + + // swConsecDroppedFrames/swFirstDroppedTime are diagnostic-only counters + // mirroring hwConsecDroppedFrames/hwFirstDroppedTime for the SW warm-up + // check above: they let the log line report how long the current + // black-frame streak has lasted, without changing drop behaviour. + swConsecDroppedFrames int + swFirstDroppedTime time.Time +} + +// extractI420fromSrc extracts I420 planar data from srcFrame into the ring +// buffer and stores a pointer in d.lastI420. Called from convertFrame() when +// outI420Enabled is true, before av_frame_unref(d.swFrame). +// Sets d.lastI420 = nil when the pixel format is not directly supported. +func (d *ffmpegDecoder) extractI420fromSrc(srcFrame *C.AVFrame) { + srcFmt := C.enum_AVPixelFormat(srcFrame.format) + if srcFmt != C.AV_PIX_FMT_YUV420P && srcFmt != C.AV_PIX_FMT_YUVJ420P && + srcFmt != C.AV_PIX_FMT_NV12 { + d.lastI420 = nil + return + } + + w := int(srcFrame.width) + h := int(srcFrame.height) + pw := (w + 1) / 2 + ph := (h + 1) / 2 + ySize := w * h + uvSize := pw * ph + + slot := &d.outI420Ring[d.outI420RingIdx] + d.outI420RingIdx ^= 1 + + if cap(slot.Y) < ySize { + slot.Y = make([]byte, ySize) + } else { + slot.Y = slot.Y[:ySize] + } + if cap(slot.U) < uvSize { + slot.U = make([]byte, uvSize) + } else { + slot.U = slot.U[:uvSize] + } + if cap(slot.V) < uvSize { + slot.V = make([]byte, uvSize) + } else { + slot.V = slot.V[:uvSize] + } + slot.YStride = w + slot.UStride = pw + slot.VStride = pw + slot.Width = w + slot.Height = h + slot.FullRange = srcFmt == C.AV_PIX_FMT_YUVJ420P || srcFrame.color_range == 2 + + if srcFmt == C.AV_PIX_FMT_YUV420P || srcFmt == C.AV_PIX_FMT_YUVJ420P { + C.grdp_copy_yuv420p_to_i420(srcFrame, + (*C.uint8_t)(unsafe.Pointer(&slot.Y[0])), + (*C.uint8_t)(unsafe.Pointer(&slot.U[0])), + (*C.uint8_t)(unsafe.Pointer(&slot.V[0])), + C.int(w), C.int(h)) + } else { + C.grdp_copy_nv12_to_i420(srcFrame, + (*C.uint8_t)(unsafe.Pointer(&slot.Y[0])), + (*C.uint8_t)(unsafe.Pointer(&slot.U[0])), + (*C.uint8_t)(unsafe.Pointer(&slot.V[0])), + C.int(w), C.int(h)) + } + + d.lastI420 = slot +} + +// extractNV12fromSrc copies native NV12 planes from srcFrame into the ring +// buffer and stores a pointer in d.lastNV12. It intentionally does not +// deinterleave chroma, so SDL2 NV12 texture uploads avoid the I420 conversion +// work required by extractI420fromSrc. +func (d *ffmpegDecoder) extractNV12fromSrc(srcFrame *C.AVFrame) { + if C.enum_AVPixelFormat(srcFrame.format) != C.AV_PIX_FMT_NV12 { + d.lastNV12 = nil + return + } + + w := int(srcFrame.width) + h := int(srcFrame.height) + uvStride := ((w + 1) / 2) * 2 + ph := (h + 1) / 2 + ySize := w * h + uvSize := uvStride * ph + + slot := &d.outNV12Ring[d.outNV12RingIdx] + d.outNV12RingIdx ^= 1 + + if cap(slot.Y) < ySize { + slot.Y = make([]byte, ySize) + } else { + slot.Y = slot.Y[:ySize] + } + if cap(slot.UV) < uvSize { + slot.UV = make([]byte, uvSize) + } else { + slot.UV = slot.UV[:uvSize] + } + slot.YStride = w + slot.UVStride = uvStride + slot.Width = w + slot.Height = h + slot.FullRange = srcFrame.color_range == 2 + + C.grdp_copy_nv12(srcFrame, + (*C.uint8_t)(unsafe.Pointer(&slot.Y[0])), + (*C.uint8_t)(unsafe.Pointer(&slot.UV[0])), + C.int(w), C.int(h)) + + d.lastNV12 = slot +} + +func newH264DecoderInternal(watchdogCh chan<- struct{}, forceSW bool, kfWaitTimeout time.Duration) rdpgfx.H264Decoder { + // Suppress FFmpeg stderr output (e.g. "[h264 @ ...] sps_id out of range"). + // grdp emits its own slog messages for H.264 recovery events. + avLogOnce.Do(func() { C.grdp_suppress_av_log() }) + + codec := C.avcodec_find_decoder(C.AV_CODEC_ID_H264) + if codec == nil { + slog.Warn("H.264: codec not found in FFmpeg") + return nil + } + + codecCtx := C.avcodec_alloc_context3(codec) + if codecCtx == nil { + return nil + } + + // alreadyOpened is set when a codec-specific path (e.g. V4L2 M2M) opens + // its own AVCodecContext before the shared avcodec_open2 call below. + alreadyOpened := false + + d := &ffmpegDecoder{ + codecCtx: codecCtx, + hwPixFmt: C.AV_PIX_FMT_NONE, + lastFmt: C.AV_PIX_FMT_NONE, + needsKeyFrame: true, // always wait for a clean IDR before feeding packets + hwNeedsZeroCheck: true, // check first NV12 output for zero-filled IOSurface + swNeedsZeroCheck: true, // check first YUV420P output for zero-fill warm-up + watchdogCh: watchdogCh, + kfWaitTimeoutVal: kfWaitTimeout, + } + + // Always enable LOW_DELAY: RDP H.264 streams are transmitted in display + // order with no B-frame reordering, so the default reorder buffer adds + // no value and (especially on VideoToolbox) makes the decoder appear + // stalled between IDRs. + C.grdp_set_low_delay(codecCtx) + + if !forceSW { + // Probe available hardware acceleration backends. + hwType := C.av_hwdevice_iterate_types(C.AV_HWDEVICE_TYPE_NONE) + for hwType != C.AV_HWDEVICE_TYPE_NONE { + var devCtx *C.AVBufferRef + if C.av_hwdevice_ctx_create(&devCtx, hwType, nil, nil, 0) == 0 { + // Find the HW pixel format for this device type. + hwPixFmt := C.enum_AVPixelFormat(C.AV_PIX_FMT_NONE) + for i := C.int(0); ; i++ { + cfg := C.avcodec_get_hw_config(codec, i) + if cfg == nil { + break + } + if cfg.device_type == hwType && + (cfg.methods&C.AV_CODEC_HW_CONFIG_METHOD_HW_DEVICE_CTX) != 0 { + hwPixFmt = cfg.pix_fmt + break + } + } + + if hwPixFmt != C.AV_PIX_FMT_NONE { + codecCtx.hw_device_ctx = C.av_buffer_ref(devCtx) + C.grdp_set_hw_pix_fmt(codecCtx, hwPixFmt) + C.grdp_set_get_format(codecCtx) + d.useHW = true + d.hwPixFmt = hwPixFmt + name := C.av_hwdevice_get_type_name(hwType) + slog.Debug("H.264: hardware acceleration enabled", "type", C.GoString(name)) + } + C.av_buffer_unref(&devCtx) + if d.useHW { + break + } + } + hwType = C.av_hwdevice_iterate_types(hwType) + } + + // If no hwdevice backend was found, try the V4L2 M2M codec (h264_v4l2m2m). + // This is a standalone FFmpeg codec that directly outputs NV12 frames in + // CPU-accessible memory and is commonly available on Linux SoCs such as + // Raspberry Pi 4/5. It is not exposed via av_hwdevice_iterate_types and + // must be probed explicitly. On macOS or when FFmpeg is built without V4L2 + // support, avcodec_find_decoder_by_name returns nil and the probe is a no-op. + if !d.useHW { + v4l2Codec := C.grdp_find_v4l2m2m() + if v4l2Codec != nil { + v4l2Ctx := C.avcodec_alloc_context3(v4l2Codec) + if v4l2Ctx != nil { + C.grdp_set_low_delay(v4l2Ctx) + if C.avcodec_open2(v4l2Ctx, v4l2Codec, nil) >= 0 { + // Replace the standard h264 context with the V4L2 M2M one. + // avcodec_free_context sets d.codecCtx to nil via its **ctx arg. + C.avcodec_free_context(&d.codecCtx) + d.codecCtx = v4l2Ctx + codecCtx = v4l2Ctx + d.useHW = true + d.hwNeedsZeroCheck = false // no zero-filled IOSurface on V4L2 + alreadyOpened = true + slog.Debug("H.264: V4L2 M2M hardware acceleration enabled") + } else { + C.avcodec_free_context(&v4l2Ctx) + } + } + } + } + } + + if !d.useHW { + if d.watchdogCh != nil { + // Main decoder switching from VideoToolbox to FFmpeg after a stall. + slog.Debug("H.264: using software decoding (SW fallback)") + } else { + // Aux decoder (h264dec2) or initial SW-only decoder — always pure SW. + slog.Debug("H.264: using software decoding") + } + // Limit the decoded picture buffer to 1 reference frame so each frame + // is emitted immediately rather than waiting for up to + // max_dec_frame_buffering (often 8) frames to accumulate. RDP H.264 + // streams use sequential P-frames that only reference the immediately + // preceding frame, so this is safe. VideoToolbox (HW path) has its + // own zero-latency output mechanism and does not need this. + codecCtx.refs = 1 + // Use slice-level threading only. Frame-level threading (the FFmpeg + // default) introduces a one-frame reorder delay that conflicts with + // AV_CODEC_FLAG_LOW_DELAY and causes each decoded frame to arrive one + // frame late — effectively doubling input latency. Slice threading + // parallelises within a single frame with no added latency, which is + // beneficial when the server encodes multiple slices per frame. + codecCtx.thread_type = C.FF_THREAD_SLICE + } + + if !alreadyOpened { + if C.avcodec_open2(codecCtx, codec, nil) < 0 { + C.avcodec_free_context(&d.codecCtx) + return nil + } + } + + d.packet = C.av_packet_alloc() + d.frame = C.av_frame_alloc() + d.swFrame = C.av_frame_alloc() + d.mapFrame = C.av_frame_alloc() + if d.packet == nil || d.frame == nil || d.swFrame == nil || d.mapFrame == nil { + d.Close() + return nil + } + + // Arm the keyframe-wait timer immediately so recovery is triggered even + // when the server sends no frames after a soft reset (e.g. static screen + // or ForceRefresh ignored by the server). If an IDR arrives first, + // Decode() cancels the timer. Decoders without a watchdog channel + // (e.g. h264dec2) are not armed here — they are managed separately. + if watchdogCh != nil { + d.kfWaitTimer = time.AfterFunc(kfWaitTimeout, func() { + d.timerBrokenReason.Store(int32(rdpgfx.H264BrokenReasonNoIDR)) + d.timerBroken.Store(true) + d.signalWatchdog() + }) + } + + runtime.SetFinalizer(d, func(dec *ffmpegDecoder) { dec.Close() }) + return d +} + +func (d *ffmpegDecoder) NeedsKeyframe() bool { + return d.needsKeyFrame +} + +func (d *ffmpegDecoder) NeedsIDR() bool { + return d.needsKeyFrame +} + +func (d *ffmpegDecoder) IsBroken() bool { + return d.broken || d.timerBroken.Load() +} + +func (d *ffmpegDecoder) BrokenReason() rdpgfx.H264BrokenReason { + if d.brokenReason != rdpgfx.H264BrokenReasonNone { + return d.brokenReason + } + return rdpgfx.H264BrokenReason(d.timerBrokenReason.Load()) +} + +func (d *ffmpegDecoder) ForceBroken(reason rdpgfx.H264BrokenReason) { + d.markBroken(reason) +} + +// markBroken sets d.broken and stops any pending background timers. +// Called from inside Decode() (decodeLoop goroutine) when a timeout fires. +func (d *ffmpegDecoder) markBroken(reason rdpgfx.H264BrokenReason) { + d.broken = true + if reason != rdpgfx.H264BrokenReasonNone { + d.brokenReason = reason + d.timerBrokenReason.Store(int32(reason)) + } + d.stopTimers() +} + +// stopTimers cancels the stall-probe and IDR-wait background timers. +func (d *ffmpegDecoder) stopTimers() { + if d.stallTimer != nil { + d.stallTimer.Stop() + d.stallTimer = nil + } + if d.kfWaitTimer != nil { + d.kfWaitTimer.Stop() + d.kfWaitTimer = nil + } +} + +// signalWatchdog sends a non-blocking signal to the GfxHandler decodeLoop so +// it calls maybeNotifyDecoderBroken even when no server frames are arriving. +func (d *ffmpegDecoder) signalWatchdog() { + if d.watchdogCh == nil { + return + } + select { + case d.watchdogCh <- struct{}{}: + default: + } +} + +// HardResetCount always returns 0 — hard resets have been removed. +// The method is kept to satisfy the rdpgfx.H264Decoder interface used by GfxHandler. +func (d *ffmpegDecoder) HardResetCount() int { + return 0 +} + +func (d *ffmpegDecoder) LastReceiveTime() time.Time { + return d.lastReceiveTime +} + +// setRegionHint specifies dirty rectangles for the next Decode call. When +// set, convertFrame will use region-aware YUV→BGRA conversion and only write +// pixels within the provided rectangles, skipping unchanged areas of the frame. +// Must be called immediately before Decode; Decode clears the hint at entry so +// it cannot accidentally apply to a later unrelated frame. +func (d *ffmpegDecoder) SetRegionHint(rects [][4]uint16) { + n := len(rects) + need := n * 4 + if cap(d.regionHint) < need { + d.regionHint = make([]C.uint16_t, need) + } else { + d.regionHint = d.regionHint[:need] + } + for i, r := range rects { + d.regionHint[i*4+0] = C.uint16_t(r[0]) + d.regionHint[i*4+1] = C.uint16_t(r[1]) + d.regionHint[i*4+2] = C.uint16_t(r[2]) + d.regionHint[i*4+3] = C.uint16_t(r[3]) + } + d.nRegionHints = C.int(n) +} + +func (d *ffmpegDecoder) Decode(h264Data []byte) (*rdpgfx.H264Frame, error) { + // Capture and clear the pending region hint immediately so that any early + // return (broken, keyframe wait, etc.) cannot leave stale hints that would + // incorrectly apply to a subsequent unrelated frame. + regHint := d.regionHint + nReg := d.nRegionHints + d.nRegionHints = 0 + + if len(h264Data) == 0 { + return nil, nil + } + if !d.outI420Enabled { + d.lastI420 = nil + } + if !d.outNV12Enabled { + d.lastNV12 = nil + } + if d.broken { + // HW decoder is unrecoverable. Stop feeding packets so no frames + // are produced; the application-level watchdog will reconnect. + return nil, nil + } + // A background timer may have fired and set timerBroken while Decode() + // was not being called (static screen → server sends no frames). + // Propagate it to broken so all downstream checks see a consistent state. + if d.timerBroken.Load() { + d.markBroken(rdpgfx.H264BrokenReason(d.timerBrokenReason.Load())) + return nil, nil + } + // Track every call, including those that return early (probe mode, keyframe + // wait, etc.). Keep the previous receive time for idle detection before we + // overwrite it with the current call timestamp. + now := time.Now() + prevReceiveTime := d.lastReceiveTime + d.lastReceiveTime = now + + // After a decoder reset we must resync with a fresh IDR from the server. + // After a SW decoder flush, wait for an IDR before resuming decoding. + // If the server never sends one within keyframeWaitLimit packets, + // attempt error-concealment decode anyway. + // FFmpeg's "[h264 @ ...] sps_id out of range" errors are suppressed at + // the av_log level (AV_LOG_FATAL) set in newH264Decoder; grdp emits its + // own slog warning instead. + // Single pass over the Annex B stream: detect IDR/SPS NAL presence. + scan := rdpgfx.ScanH264Packet(h264Data) + + if d.needsKeyFrame { + if !scan.HasKeyFrame { + d.keyframeWaitCount++ + if d.keyframeWaitCount == 1 { + d.keyframeWaitStart = time.Now() + if d.useHW { + slog.Debug("H.264: HW decoder waiting for IDR") + // kfWaitTimer was armed at decoder creation; only start a new + // one here if the decoder was created without a watchdog channel + // (no timer was armed at creation time). + if d.kfWaitTimer == nil { + d.kfWaitTimer = time.AfterFunc(d.kfWaitTimeoutVal, func() { + d.timerBrokenReason.Store(int32(rdpgfx.H264BrokenReasonNoIDR)) + d.timerBroken.Store(true) + d.signalWatchdog() + }) + } + } + } else if d.keyframeWaitCount%30 == 0 { + slog.Debug("H.264: still waiting for IDR", + "waited", d.keyframeWaitCount, + "waitedFor", time.Since(d.keyframeWaitStart).Round(time.Millisecond)) + } + kfTimeout := d.kfWaitTimeoutVal // use per-decoder limit (HW: 15 s, SW fallback: 8 s) + if !d.useHW && d.watchdogCh == nil { + // Detached aux SW decoder (h264dec2): shorter wait so it is + // torn down and recreated quickly on the next stream2 IDR. + kfTimeout = keyframeWaitTimeoutSW // 5 s + } + waitedTooLong := !d.keyframeWaitStart.IsZero() && + time.Since(d.keyframeWaitStart) >= kfTimeout + if d.keyframeWaitCount >= keyframeWaitLimit || waitedTooLong { + if d.useHW || d.watchdogCh != nil { + // HW decoders (e.g. VideoToolbox) and watchdog-armed SW + // decoders (main decoder SW fallback) cannot recover + // without a proper IDR. Mark broken so the recovery + // chain can escalate. For the SW fallback case, error- + // concealment on P-frames without reference frames always + // fails (avcodec_send_packet returns EINVAL) and would + // only produce a spurious WARN and an immediate reconnect. + slog.Debug("H.264: no IDR received, marking broken", + "hw", d.useHW, + "waited", d.keyframeWaitCount, + "waitedFor", time.Since(d.keyframeWaitStart).Round(time.Millisecond)) + d.markBroken(rdpgfx.H264BrokenReasonNoIDR) + return nil, nil + } + slog.Debug("H.264: aux SW decoder: no IDR received, attempting error-concealment", + "waited", d.keyframeWaitCount, + "waitedFor", time.Since(d.keyframeWaitStart).Round(time.Millisecond)) + d.needsKeyFrame = false + d.keyframeWaitCount = 0 + d.keyframeWaitStart = time.Time{} + d.proceededWithoutKeyframe = true + // fall through and attempt SW error-concealment decode + } else { + return nil, nil // drop P-frames while waiting + } + } else { + waitedFor := time.Duration(0) + if !d.keyframeWaitStart.IsZero() { + waitedFor = time.Since(d.keyframeWaitStart).Round(time.Millisecond) + } + slog.Debug("H.264: IDR received, resuming decode", + "hw", d.useHW, "waitedFor", waitedFor) + d.needsKeyFrame = false + d.keyframeWaitCount = 0 + d.keyframeWaitStart = time.Time{} + // IDR received — cancel the background wait timer. + if d.kfWaitTimer != nil { + d.kfWaitTimer.Stop() + d.kfWaitTimer = nil + } + } + } + + // If we previously proceeded without a keyframe (error-concealment path) + // and the server has now sent a proper IDR, the decoder is back to a clean + // state — clear the flag so a future send failure is not misattributed to + // the (long-past) keyframe wait exhaustion. + if d.proceededWithoutKeyframe && scan.HasKeyFrame { + d.proceededWithoutKeyframe = false + } + + // VideoToolbox sometimes returns a zero-filled IOSurface on the first + // frame after any IDR — not only after decoder creation or flush — because + // the hardware pipeline must drain and reset its reference frames before it + // can produce the new intra frame. Re-arm the zero-check whenever we + // receive an IDR so that convertFrame discards any spurious green frame that + // VideoToolbox outputs during that transition. + if d.useHW && scan.HasKeyFrame { + d.hwNeedsZeroCheck = true + } + + // Time-based stall detection for the HW decoder. + // + // hwReady=false: decoder has never produced a frame. If it keeps receiving + // packets without ever outputting anything, the VideoToolbox session failed + // to initialise — mark broken so the soft-reset/reconnect path fires. + // + // hwReady=true: decoder was working. VideoToolbox legitimately stalls for + // several seconds when processing an IDR / scene-change keyframe (it must + // flush its internal reference pipeline before it can resume output). + // Firing broken on these stalls causes unnecessary soft-reset loops. We + // apply avcHWReadyFreezeThreshold here as a pre-flight guard: if the + // decoder has been silent for longer than the threshold we mark it broken + // and return *without* calling avcodec_send_packet. This is critical + // because on macOS VideoToolbox the CGo call itself permanently blocks + // after ~5.75 s of stall, permanently hanging the decodeLoop goroutine. + // + // False-positive guard: if the RDP server itself was idle (no packets sent + // for at least avcHWReadyFreezeThreshold), the elapsed time since + // lastSuccessTime reflects server silence, not a VideoToolbox deadlock. + // In that case we reset the stall clock so the threshold applies only to + // periods where packets were actually flowing into the decoder. + if d.useHW && !d.hwReady && !d.hwFirstSendTime.IsZero() { + if stalledFor := time.Since(d.hwFirstSendTime); stalledFor >= avcFreezeThreshold { + slog.Warn("H.264: HW decoder failed to produce first frame, marking broken", + "stalledFor", stalledFor, "hwSentCount", d.hwSentCount) + d.markBroken(rdpgfx.H264BrokenReasonInitFailure) + return nil, nil + } + } + if d.useHW && d.hwReady && !d.lastSuccessTime.IsZero() { + // Early probe: the null-frame count detector may have set stallProbeStart + // before stalledFor reached readyThreshold. Handle it here so we skip + // avcodec_send_packet during the probe window even while stalledFor is + // still below the 7-second CGo-safe threshold. + if !d.stallProbeStart.IsZero() { + readyThreshold := avcHWReadyFreezeThreshold + if d.hwSentCount < avcHWEarlyFrameLimit { + readyThreshold = avcHWEarlyFreezeThreshold + } + if stalledFor := time.Since(d.lastSuccessTime); stalledFor < readyThreshold { + // Probe active but main threshold not yet crossed. Try to drain + // a frame that VT may have buffered; if found the stall was + // transient and we resume normally. + if C.avcodec_receive_frame(d.codecCtx, d.frame) >= 0 { + C.av_frame_unref(d.frame) + d.lastSuccessTime = time.Now() + d.hwConsecNullFrames = 0 + d.stallProbeStart = time.Time{} + if d.stallTimer != nil { + d.stallTimer.Stop() + d.stallTimer = nil + } + slog.Debug("H.264: HW decoder recovered during early probe (drain found frame)", + "hwSentCount", d.hwSentCount) + // Fall through to send the current packet normally. + } else if probedFor := time.Since(d.stallProbeStart); probedFor >= avcHWRecoveryWindow { + slog.Debug("H.264: HW decoder early-probe timed out, marking broken", + "probedFor", probedFor.Round(time.Second), + "frozenFor", stalledFor.Round(time.Second), + "hwSentCount", d.hwSentCount) + d.markBroken(rdpgfx.H264BrokenReasonHWStall) + return nil, nil + } else { + // Still inside probe window: skip send_packet to avoid + // feeding the stalled VT pipeline. + return nil, nil + } + } + // else: stalledFor >= readyThreshold — fall through to the + // threshold-based block below which also handles the probe. + } + } + if d.useHW && d.hwReady && !d.lastSuccessTime.IsZero() { + readyThreshold := avcHWReadyFreezeThreshold + if d.hwSentCount < avcHWEarlyFrameLimit { + readyThreshold = avcHWEarlyFreezeThreshold + } + if stalledFor := time.Since(d.lastSuccessTime); stalledFor >= readyThreshold { + // If no packet had arrived since the previous Decode() call during + // the apparent stall, the server was simply idle (e.g. screen was + // static). Reset the stall clock so we don't misfire on the first + // packet after a server-side pause. + if prevReceiveTime.IsZero() || now.Sub(prevReceiveTime) >= readyThreshold { + slog.Debug("H.264: HW decoder stall clock reset (server was idle)", + "idleFor", stalledFor, "hwSentCount", d.hwSentCount) + d.lastSuccessTime = now + d.stallProbeStart = time.Time{} + if d.stallTimer != nil { + d.stallTimer.Stop() + d.stallTimer = nil + } + } else { + // Probe for pending output that VideoToolbox may be about to + // produce. VT legitimately stalls for several seconds at a + // GOP/IDR boundary while it flushes its reference pipeline; + // immediately marking broken would cause an unnecessary + // soft-reset loop followed by a ForceRefresh that the server + // may not honour with a timely IDR. + // + // avcodec_receive_frame is non-blocking and safe to call + // without a preceding send_packet. If a frame emerges VT was + // just slow but is still healthy — reset the stall clock and + // let the current packet be sent normally below. + if C.avcodec_receive_frame(d.codecCtx, d.frame) >= 0 { + C.av_frame_unref(d.frame) + d.lastSuccessTime = time.Now() + d.stallProbeStart = time.Time{} + // Stall resolved — cancel the background probe timer. + if d.stallTimer != nil { + d.stallTimer.Stop() + d.stallTimer = nil + } + slog.Debug("H.264: HW decoder stall clock reset (drain found pending frame)", + "hadBeenSilentFor", stalledFor, "hwSentCount", d.hwSentCount) + // Fall through to send the current packet normally. + } else { + // No output yet. Enter / stay in recovery-probe window. + if d.stallProbeStart.IsZero() { + d.stallProbeStart = now + slog.Debug("H.264: HW decoder stall detected, probing for recovery", + "frozenFor", stalledFor.Round(time.Millisecond), + "hwSentCount", d.hwSentCount) + // Start a background timer so the probe window expires + // even when the server sends no more frames. + d.stallTimer = time.AfterFunc(avcHWRecoveryWindow, func() { + d.timerBrokenReason.Store(int32(rdpgfx.H264BrokenReasonHWStall)) + d.timerBroken.Store(true) + d.signalWatchdog() + }) + } else if probedFor := time.Since(d.stallProbeStart); probedFor >= avcHWRecoveryWindow { + slog.Debug("H.264: HW decoder recovery probe timed out, marking broken", + "totalFrozen", stalledFor.Round(time.Second), + "probedFor", probedFor.Round(time.Second), + "hwSentCount", d.hwSentCount) + d.markBroken(rdpgfx.H264BrokenReasonHWStall) + return nil, nil + } + // Still in recovery window: skip send_packet to avoid the + // ~5.75 s CGo deadlock and wait for VT to resume. + return nil, nil + } + } + } else if !d.stallProbeStart.IsZero() { + // Stall resolved (lastSuccessTime updated by normal frame output). + slog.Debug("H.264: HW decoder recovered from stall", + "probedFor", time.Since(d.stallProbeStart).Round(time.Millisecond)) + d.stallProbeStart = time.Time{} + // Cancel the background probe timer — VT recovered on its own. + if d.stallTimer != nil { + d.stallTimer.Stop() + d.stallTimer = nil + } + } + } + + // Pass the Go slice's backing array directly to avcodec_send_packet + // instead of allocating + copying via C.CBytes for every packet. + // FFmpeg copies the buffer internally for non-refcounted packets, so the + // memory only needs to remain valid for the duration of the C call — + // runtime.KeepAlive guarantees this. + d.packet.data = (*C.uint8_t)(unsafe.Pointer(&h264Data[0])) + d.packet.size = C.int(len(h264Data)) + + // Count packets sent to HW decoder (for init timeout tracking). + if d.useHW { + d.hwSentCount++ + hwNow := time.Now() + if d.hwSentCount == 1 { + d.hwFirstSendTime = hwNow + } + d.lastSendTime = hwNow + } + + sendStart := time.Now() + ret := C.avcodec_send_packet(d.codecCtx, d.packet) + sendNs := time.Since(sendStart).Nanoseconds() + // Make sure the Go-managed h264Data backing array is not collected or + // moved while FFmpeg is reading from it inside the C call above. + runtime.KeepAlive(h264Data) + // Drop the Go pointer from the AVPacket immediately so a subsequent + // avcodec_* call can't dereference stale memory. + d.packet.data = nil + d.packet.size = 0 + if ret < 0 { + // Both HW and SW: flush the decoder pipeline and wait for a fresh IDR. + // Reset the HW stall-timer so it starts fresh after the IDR arrives, + // not from before this failed send attempt. + slog.Debug("H.264: avcodec_send_packet failed, flushing decoder to recover", + "hw", d.useHW, "err", int(ret)) + C.avcodec_flush_buffers(d.codecCtx) + prev := d.proceededWithoutKeyframe + d.needsKeyFrame = true + d.keyframeWaitCount = 0 + d.keyframeWaitStart = time.Time{} + d.proceededWithoutKeyframe = false + if d.useHW { + d.hwFirstSendTime = time.Time{} // restart stall clock after IDR + d.hwSentCount = 0 + d.hwConsecNullFrames = 0 + d.hwConsecDroppedFrames = 0 + d.hwFirstDroppedTime = time.Time{} + d.hwNeedsZeroCheck = true // re-check for zero-filled IOSurface after flush + if !d.hwReady && prev { + // We gave up waiting for an IDR and tried a P-frame anyway, and + // VideoToolbox rejected it. There is no further recovery possible + // for this decoder context — mark broken so the soft-reset / + // reconnect chain can proceed. + slog.Warn("H.264: HW decoder rejected packet after keyframe wait exhaustion, marking broken", + "err", int(ret)) + d.markBroken(rdpgfx.H264BrokenReasonNoIDR) + } + } else { + // Re-arm the SW zero-check: libavcodec similarly outputs zero frames + // on the first packet after a flush, especially when primed with a + // stale cached IDR. + d.swNeedsZeroCheck = true + if prev { + // SW decoder: error-concealment attempt (proceededWithoutKeyframe) + // failed — avcodec_send_packet rejected the P-frame. Without + // marking broken the decoder would loop: wait 900 frames → try → + // fail → reset → wait 900 frames → ... Mark broken so the + // soft-reset / reconnect chain can escalate instead. + slog.Warn("H.264: SW decoder rejected packet after keyframe wait exhaustion, marking broken", + "err", int(ret)) + d.markBroken(rdpgfx.H264BrokenReasonNoIDR) + } + } + return nil, nil + } + + // Receive decoded frame(s); keep the last one. + var result *rdpgfx.H264Frame + var recvNs, convertNs, transferNs, maxConvNs int64 + for { + recvStart := time.Now() + ret = C.avcodec_receive_frame(d.codecCtx, d.frame) + recvNs += time.Since(recvStart).Nanoseconds() + if ret < 0 { + break // EAGAIN (need more input) or EOF + } + convStart := time.Now() + f, tNs, err := d.convertFrame(regHint, nReg) + dur := time.Since(convStart).Nanoseconds() + convertNs += dur + transferNs += tNs + if dur > maxConvNs { + maxConvNs = dur + } + C.av_frame_unref(d.frame) + if err != nil { + return nil, err + } + result = f + } + + // I420/NV12 fast paths return nil for the BGRA frame but still represent a + // successfully decoded frame. Count them as success for health tracking. + gotFrame := result != nil || d.lastI420 != nil || d.lastNV12 != nil + if gotFrame { + d.lastSuccessTime = time.Now() + d.hwConsecNullFrames = 0 + if d.useHW { + // A dropped frame (warm-up/zero-fill/low-chroma junk — see + // convertFrame) still counts as "gotFrame" above, which keeps + // resetting the null-frame/lastSuccessTime stall detectors. If + // VideoToolbox never produces anything but junk, those detectors + // would otherwise never fire and the app would freeze on a black + // frame forever. Track dropped-frame streaks independently and + // escalate to broken/HWStall once the streak is clearly abnormal. + if result != nil && result.Dropped { + if d.hwConsecDroppedFrames == 0 { + d.hwFirstDroppedTime = time.Now() + } + d.hwConsecDroppedFrames++ + droppedFor := time.Since(d.hwFirstDroppedTime) + if d.hwConsecDroppedFrames >= avcHWConsecDroppedFrameLimit || + (d.hwConsecDroppedFrames >= avcHWConsecDroppedFrameMinCount && + droppedFor >= avcHWDroppedFrameStallThreshold) { + slog.Warn("H.264: HW decoder producing only warm-up/junk frames, marking broken", + "consecDroppedFrames", d.hwConsecDroppedFrames, + "droppedFor", droppedFor.Round(time.Millisecond), + "hwSentCount", d.hwSentCount) + d.markBroken(rdpgfx.H264BrokenReasonHWStall) + return nil, nil + } + } else { + d.hwConsecDroppedFrames = 0 + d.hwFirstDroppedTime = time.Time{} + } + if !d.hwReady { + slog.Debug("H.264: HW decoder produced first frame", + "hwSentCount", d.hwSentCount) + } + d.hwReady = true + + // Aggregate per-frame timing for the HW path. + d.profFrames++ + d.profSendNs += sendNs + d.profRecvNs += recvNs + d.profConvertNs += convertNs + d.profTransferNs += transferNs + if maxConvNs > d.profMaxConvNs { + d.profMaxConvNs = maxConvNs + } + if sendNs > d.profMaxSendNs { + d.profMaxSendNs = sendNs + } + if recvNs > d.profMaxRecvNs { + d.profMaxRecvNs = recvNs + } + if d.profFrames >= profileWindow { + n := int64(d.profFrames) + slog.Debug("H.264: HW decode timing", + "frames", d.profFrames, + "avgSendUs", d.profSendNs/n/1000, + "avgRecvUs", d.profRecvNs/n/1000, + "avgConvertUs", d.profConvertNs/n/1000, + "avgTransferUs", d.profTransferNs/n/1000, + "maxSendUs", d.profMaxSendNs/1000, + "maxRecvUs", d.profMaxRecvNs/1000, + "maxConvertUs", d.profMaxConvNs/1000) + d.profFrames = 0 + d.profSendNs = 0 + d.profRecvNs = 0 + d.profConvertNs = 0 + d.profTransferNs = 0 + d.profMaxConvNs = 0 + d.profMaxSendNs = 0 + d.profMaxRecvNs = 0 + } + } + } else { // !gotFrame + if d.useHW && d.hwReady { + stalledFor := time.Since(d.lastSuccessTime) + d.hwConsecNullFrames++ + slog.Debug("H.264: HW null frame", "frozenFor", stalledFor, + "hwSentCount", d.hwSentCount) + // Stall probe: trigger a probe if we accumulate many consecutive + // null frames before the 7-second CGo-safe threshold. Two tiers: + // • Early window (hwSentCount < avcHWEarlyFrameLimit): use + // avcHWNullFrameStallLimit (25). Reduces visible freeze from + // ~10 s to ~4 s for a genuine VT stall at session start. + // • Mid-session (hwSentCount >= avcHWEarlyFrameLimit): use + // avcHWMidSessionNullFrameLimit (30 ≈ 1 s at 30 fps). + // Normal GOP boundaries produce ≤25 null frames so there is + // minimal headroom; genuine stalls persist for hundreds of + // null frames. Reduces visible freeze from ~2.5 s to ~1 s. + // + // IDR-flush suppression: when VideoToolbox just received a new IDR + // (hwNeedsZeroCheck=true), null frames are part of its normal pipeline + // flush — it must drain its internal reference frames before outputting + // the new intra picture. Triggering a stall probe on these IDR-induced + // null frames causes premature SW fallback and an unnecessary reconnect: + // observed logs show VT recovering naturally within ~1-2 s of the IDR, + // and the post-reconnect VT session exhibits the same null-frame burst + // (which resolves on its own). Suppress the count-based probe while + // hwNeedsZeroCheck is true; the 7-second safety valve below remains + // the backstop for genuine stalls that do not self-resolve. + earlyStall := d.hwSentCount < avcHWEarlyFrameLimit && + d.hwConsecNullFrames >= avcHWNullFrameStallLimit && + !d.hwFirstSendTime.IsZero() && + time.Since(d.hwFirstSendTime) >= avcHWEarlyStallMinElapsed + midStall := d.hwSentCount >= avcHWEarlyFrameLimit && + d.hwConsecNullFrames >= avcHWMidSessionNullFrameLimit && + !d.hwNeedsZeroCheck // IDR-flush null frames: let VT recover naturally + if (earlyStall || midStall) && d.stallProbeStart.IsZero() { + slog.Debug("H.264: HW decoder stall detected (null frame count), entering probe", + "consecNullFrames", d.hwConsecNullFrames, + "frozenFor", stalledFor.Round(time.Millisecond), + "hwSentCount", d.hwSentCount) + d.stallProbeStart = time.Now() + d.stallTimer = time.AfterFunc(avcHWRecoveryWindow, func() { + d.timerBrokenReason.Store(int32(rdpgfx.H264BrokenReasonHWStall)) + d.timerBroken.Store(true) + d.signalWatchdog() + }) + } + // Safety valve: if the pre-flight probe window is NOT active and + // the decoder has been silent past the threshold, VideoToolbox is + // genuinely stuck. In probe mode the pre-flight block (above) is + // responsible for declaring the decoder broken — the safety valve + // must not interfere with the probe window countdown. + if d.stallProbeStart.IsZero() && stalledFor >= avcHWReadyFreezeThreshold { + slog.Warn("H.264: HW decoder stall timeout (safety valve), marking broken", + "frozenFor", stalledFor, "hwSentCount", d.hwSentCount) + d.markBroken(rdpgfx.H264BrokenReasonHWStall) + } + } + } + return result, nil +} + +// DecodeWithI420 implements the rdpgfx.I420Decoder interface. It decodes H.264 NAL +// data and returns both a BGRA frame (for the surface backing store) and an +// optional I420 frame for GPU-accelerated rendering via SDL2 IYUV textures. +// The I420 frame is nil when the decoder's pixel format is not directly +// supported (e.g. swscale paths that have already consumed the source frame +// before we could extract planar data, or hardware-decoded frames whose +// transfer format is not YUV420P or NV12). Callers must fall back to BGRA +// rendering when I420 is nil. +func (d *ffmpegDecoder) DecodeWithI420(h264Data []byte) (*rdpgfx.H264Frame, *rdpgfx.H264FrameI420, error) { + d.outI420Enabled = true + d.lastI420 = nil + frame, err := d.Decode(h264Data) + d.outI420Enabled = false + return frame, d.lastI420, err +} + +// LastI420 returns the I420 frame produced during the most recent Decode, +// DecodeWithI420, or DecodeWithNV12 call. For DecodeWithNV12, this is +// non-nil when the source format was YUV420P/YUVJ420P (software decoder) +// rather than NV12, allowing callers to refresh the AVC444 Y cache even +// when no native NV12 planes were available. Must be called from the same +// goroutine as Decode; the returned pointer is valid until the next call. +func (d *ffmpegDecoder) LastI420() *rdpgfx.H264FrameI420 { + return d.lastI420 +} + +// DecodeWithNV12 implements the rdpgfx.NV12Decoder interface. It decodes H.264 NAL +// data and returns native NV12 output when FFmpeg produces NV12, avoiding the +// extra NV12->I420 deinterleave used by DecodeWithI420. +func (d *ffmpegDecoder) DecodeWithNV12(h264Data []byte) (*rdpgfx.H264Frame, *rdpgfx.H264FrameNV12, error) { + d.outNV12Enabled = true + d.lastNV12 = nil + frame, err := d.Decode(h264Data) + d.outNV12Enabled = false + return frame, d.lastNV12, err +} + +func (d *ffmpegDecoder) convertFrame(regionHint []C.uint16_t, nRegions C.int) (*rdpgfx.H264Frame, int64, error) { + srcFrame := d.frame + var transferNs int64 + usedMapFrame := false + + // Transfer from GPU to CPU memory if using hardware acceleration. + if d.useHW && d.frame.format == C.int(d.hwPixFmt) { + tStart := time.Now() + // Prefer zero-copy CPU mapping (av_hwframe_map) over a copy + // (av_hwframe_transfer_data). VideoToolbox on macOS stores decoded + // frames in IOSurface-backed shared memory, so mapping is supported + // and avoids a full GPU→RAM copy of the pixel data. + // hwMapSupported: 0=unknown (first frame), 1=ok, -1=unsupported. + if d.hwMapSupported >= 0 { + C.av_frame_unref(d.mapFrame) + if ret := C.grdp_hwframe_map(d.mapFrame, d.frame); ret >= 0 { + d.hwMapSupported = 1 + srcFrame = d.mapFrame + usedMapFrame = true + } else if d.hwMapSupported == 0 { + // First attempt failed; mark unsupported and fall through. + d.hwMapSupported = -1 + } + } + if !usedMapFrame { + ret := C.av_hwframe_transfer_data(d.swFrame, d.frame, 0) + transferNs = time.Since(tStart).Nanoseconds() + if ret < 0 { + return nil, transferNs, fmt.Errorf("av_hwframe_transfer_data: error %d", int(ret)) + } + srcFrame = d.swFrame + } else { + transferNs = time.Since(tStart).Nanoseconds() + } + } + + w := srcFrame.width + h := srcFrame.height + srcFmt := C.enum_AVPixelFormat(srcFrame.format) + + // Fast path for SDL2 NV12 texture upload. VideoToolbox usually transfers + // hardware-decoded H.264 frames as NV12, so keeping the interleaved UV plane + // intact avoids the chroma deinterleave required by I420. + if d.outNV12Enabled && srcFmt == C.AV_PIX_FMT_NV12 { + // Apply the same zero-filled IOSurface check as the BGRA NV12 path + // (see comment near hwNeedsZeroCheck below). The NV12 fast path + // previously bypassed this check, allowing zero UV (→ green) frames to + // reach the NV12 callback. + if d.useHW && d.hwNeedsZeroCheck { + var sy, su, sv C.uint8_t + C.grdp_sample_nv12(srcFrame, &sy, &su, &sv) + drop := false + if su == 0 && sv == 0 { + drop = true + slog.Debug("H.264: dropping zero-UV HW frame in NV12 path (IOSurface not ready)", + "Y", int(sy)) + } else if sy == 0 && su >= 124 && su <= 132 && sv >= 124 && sv <= 132 && + C.grdp_is_warmup_nv12(srcFrame, 4) != 0 { + drop = true + // Diagnostic-only fields: streak duration/count come from the + // caller's tracking (set on the previous drop, since this call + // happens before the caller increments the counters for THIS + // frame), and the 3x3 Y grid confirms whether the whole + // sampled area is really black or just the single centre + // pixel checked above. Logged on every Nth drop (not every + // frame) to avoid flooding, since a real bug reproduction may + // run for many seconds at high frame rate. + streakElapsed := time.Duration(0) + if !d.hwFirstDroppedTime.IsZero() { + streakElapsed = time.Since(d.hwFirstDroppedTime) + } + if d.hwConsecDroppedFrames%10 == 0 { + var grid [9]C.uint8_t + C.grdp_debug_sample_y_grid(srcFrame, &grid[0]) + ys := make([]int, 9) + for i, v := range grid { + ys[i] = int(v) + } + slog.Debug("H.264: dropping black warm-up HW frame in NV12 path", + "Y", int(sy), "U", int(su), "V", int(sv), + "streakCount", d.hwConsecDroppedFrames, + "streakElapsed", streakElapsed.Round(time.Millisecond), + "yGrid", ys) + } else { + slog.Debug("H.264: dropping black warm-up HW frame in NV12 path", + "Y", int(sy), "U", int(su), "V", int(sv), + "streakCount", d.hwConsecDroppedFrames, + "streakElapsed", streakElapsed.Round(time.Millisecond)) + } + } + if drop { + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + d.hwNeedsZeroCheck = false + } + // Permanent guard: the single-pixel check above only runs while + // hwNeedsZeroCheck is armed, but stale IDR priming in the SW fallback + // (and rare HW corruption) can produce green-monochrome frames later in + // the session. Sample a 3x3 grid and drop frames whose chroma has + // collapsed to ~0 across most of the frame. + if C.grdp_is_low_chroma_nv12(srcFrame, lowChromaThreshold, 6) != 0 { + slog.Debug("H.264: dropping low-chroma NV12 frame (green-monochrome corruption)") + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + d.extractNV12fromSrc(srcFrame) + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return nil, transferNs, nil + } + + // Fast path: when I420 output is requested and the source pixel format is + // directly convertible to I420 (YUV420P, YUVJ420P, NV12), skip the + // YUV→BGRA conversion entirely. The SDL2 IYUV texture render path does + // not need BGRA; eliminating the conversion saves roughly w*h*4 bytes of + // CPU writes per frame (≈8 MB at 1920×1080). + // Trade-off: blitToSurface will not be called for this frame, so the + // RDPGFX surface backing store will not reflect the H.264 content. + // SurfaceToSurface reads from this surface will see stale data, but in + // practice H.264-decoded surfaces are destination-only in normal sessions. + if d.outI420Enabled { + if srcFmt == C.AV_PIX_FMT_YUV420P || srcFmt == C.AV_PIX_FMT_YUVJ420P || + srcFmt == C.AV_PIX_FMT_NV12 { + // Apply the zero-filled IOSurface check before extracting I420. + // The I420 fast path previously bypassed hwNeedsZeroCheck entirely: + // VT could output Y≠0, UV=0 (partially initialised IOSurface) which + // isNullYUVFrame would not catch (only Y=0 && UV=0 triggers it), + // causing a bright-green frame. Additionally, even an all-zero + // frame (Y=0, UV=0) would poison the AVC444 Y cache via + // updateAVC444YCache(), producing green LC=2 combine artifacts for + // the next 500 ms window. Check UV here and drop the frame if zero, + // matching the BGRA NV12 path behaviour. + if srcFmt == C.AV_PIX_FMT_NV12 && d.useHW && d.hwNeedsZeroCheck { + var sy, su, sv C.uint8_t + C.grdp_sample_nv12(srcFrame, &sy, &su, &sv) + drop := false + if su == 0 && sv == 0 { + drop = true + slog.Debug("H.264: dropping zero-UV HW frame in I420 path (IOSurface not ready)", + "Y", int(sy)) + } else if sy == 0 && su >= 124 && su <= 132 && sv >= 124 && sv <= 132 && + C.grdp_is_warmup_nv12(srcFrame, 4) != 0 { + drop = true + slog.Debug("H.264: dropping black warm-up HW frame in I420 path", + "Y", int(sy), "U", int(su), "V", int(sv)) + } + if drop { + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + // Return Dropped so Decode() counts this as success (health + // tracking stays correct) while signalling grdp to skip the + // I420 callback and Y-cache update. + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + d.hwNeedsZeroCheck = false + } + // Permanent low-chroma guard for the I420 fast path, matching the + // NV12 fast path above. This protects the SDL2 IYUV texture and the + // AVC444 Y cache from green-monochrome corruption. + if srcFmt == C.AV_PIX_FMT_NV12 { + if C.grdp_is_low_chroma_nv12(srcFrame, lowChromaThreshold, 6) != 0 { + slog.Debug("H.264: dropping low-chroma NV12 frame in I420 path (green-monochrome corruption)") + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + } else if srcFmt == C.AV_PIX_FMT_YUV420P || srcFmt == C.AV_PIX_FMT_YUVJ420P { + if C.grdp_is_low_chroma_yuv420p(srcFrame, lowChromaThreshold, 6) != 0 { + slog.Debug("H.264: dropping low-chroma YUV420P frame in I420 path (green-monochrome corruption)") + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + } + d.extractI420fromSrc(srcFrame) + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return nil, transferNs, nil + } + } + + outSize := int(w) * int(h) * 4 + // Borrow the next ring buffer instead of allocating fresh. At 1920×1080 + // this avoids an 8MB allocation every frame. + out := d.outRing[d.outRingIdx] + if cap(out) < outSize { + out = make([]byte, outSize) + } else { + out = out[:outSize] + } + d.outRing[d.outRingIdx] = out + d.outRingIdx ^= 1 + + // For planar YUV420P (both limited- and full-range variants), use our own + // BT.601 conversion instead of swscale on ARM64. swscale has no + // accelerated colorspace-conversion path for yuv420p→bgra on ARM64 and + // its non-accelerated fallback ignores sws_setColorspaceDetails, + // producing a strong green cast. On x86_64 swscale is both correct and + // significantly faster (SIMD-accelerated), so we route through swscale + // there and only fall back to the hand-written loop on ARM64. + // + // For NV12 (VideoToolbox HW transfer output) on ARM64, bypass swscale for + // the same reason: the non-accelerated ARM64 path ignores + // sws_setColorspaceDetails and produces a green cast on zero-filled frames. + // + // Exception: when dirty-region hints are provided, always use the + // hand-written region-aware functions even on x86_64. swscale has no + // partial-frame API, so it would convert the full frame unconditionally. + // For typical RDP partial-screen updates (cursors, small windows) the + // scalar BT.601 loop over only the dirty pixels is significantly faster + // than running swscale over the entire frame. + haveRegions := nRegions > 0 && len(regionHint) > 0 + var convErr error + switch { + case (srcFmt == C.AV_PIX_FMT_YUV420P || srcFmt == C.AV_PIX_FMT_YUVJ420P) && (!useSwscale || haveRegions): + fullRange := C.int(0) + if srcFmt == C.AV_PIX_FMT_YUVJ420P || srcFrame.color_range == 2 { + fullRange = 1 + } + // Sample the centre pixel for diagnostic logging and SW zero-frame checks. + // For SW decoder, always sample: zero-UV is checked on every frame. + needSample := d.hwFrameCount < 3 || !d.useHW + var sy, su, sv C.uint8_t + if needSample { + C.grdp_sample_yuv(srcFrame, &sy, &su, &sv) + } + // Log the centre-pixel YUV values for the first few frames so we + // can distinguish H.264 decode corruption from colour-conversion bugs. + if d.hwFrameCount < 3 || (!d.useHW && d.swFrameCount < 3) { + slog.Debug("H.264: frame sample (yuv420p)", + "hw", d.useHW, + "frame", d.hwFrameCount, + "fmt", int(srcFmt), + "colorRange", int(srcFrame.color_range), + "fullRange", int(fullRange), + "Y", int(sy), "U", int(su), "V", int(sv), + "w", int(w), "h", int(h)) + if d.useHW { + d.hwFrameCount++ + } else { + d.swFrameCount++ + } + } + // Drop corrupted/uninitialised frames from the SW decoder. + // + // Zero-UV check (permanent): U=0 and V=0 simultaneously never occurs in + // real BT.601/BT.709 content — valid chroma always centres on 128. This + // pattern indicates a reference-frame mismatch (stale-IDR priming with + // live P-frames that reference a different state) or an uninitialised + // buffer. BT.601 conversion of (Y=0, U=0, V=0) → BGRA(0,135,0,255) + // is a full-screen bright-green frame. Apply permanently; cost is one + // pixel sample per frame which is negligible. + // + // Near-zero UV check: stale-IDR priming sometimes produces chroma that + // is not exactly 0/0 but still extremely low (e.g. U=0, V=2). Valid + // content, even dark scenes, keeps chroma centred near 128; both planes + // collapsing to ~0 is a reliable sign of decoder corruption and renders + // as a full-screen green frame (BT.601 of Cb≈0,Cr≈0 → BGRA(0,~135,0)). + // Drop when U and V are both abnormally low regardless of luma: corrupt + // SW-fallback frames have been observed at Y≈40, U≈60, V≈42, which is + // well below the healthy desktop-content range (U≈101–122, V≈88–122) + // but above the old threshold of 24. lowChromaThreshold (72) catches + // these green-monochrome frames without affecting valid content. + // + // Warm-up black check (gated by swNeedsZeroCheck): Y=0, U≈128, V≈128 + // with all luma near zero is libavcodec's initial black output before + // the pipeline is ready. Only checked around decoder creation/flush. + if !d.useHW { + if su == 0 && sv == 0 { + slog.Debug("H.264: dropping zero-UV SW frame (reference mismatch or uninitialised)", + "Y", int(sy)) + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + if su < lowChromaThreshold && sv < lowChromaThreshold { + slog.Debug("H.264: dropping near-zero-UV SW frame (stale IDR prime?)", + "Y", int(sy), "U", int(su), "V", int(sv)) + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + // The centre-pixel check above can miss corruption that leaves the + // centre valid while the rest of the frame collapses to near-zero + // chroma. Sample a 3x3 grid and drop if most samples are abnormally + // low — this catches the green-monochrome frames produced by stale + // IDR priming in the SW fallback decoder. + if C.grdp_is_low_chroma_yuv420p(srcFrame, lowChromaThreshold, 6) != 0 { + slog.Debug("H.264: dropping low-chroma YUV420P frame in BGRA path (green-monochrome corruption)") + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + if d.swNeedsZeroCheck { + if sy == 0 && su >= 124 && su <= 132 && sv >= 124 && sv <= 132 && + C.grdp_is_warmup_nv12(srcFrame, 4) != 0 { + if d.swConsecDroppedFrames == 0 { + d.swFirstDroppedTime = time.Now() + } + d.swConsecDroppedFrames++ + streakElapsed := time.Since(d.swFirstDroppedTime) + if d.swConsecDroppedFrames%10 == 1 { + var grid [9]C.uint8_t + C.grdp_debug_sample_y_grid(srcFrame, &grid[0]) + ys := make([]int, 9) + for i, v := range grid { + ys[i] = int(v) + } + slog.Debug("H.264: dropping black warm-up SW frame", + "Y", int(sy), "U", int(su), "V", int(sv), + "streakCount", d.swConsecDroppedFrames, + "streakElapsed", streakElapsed.Round(time.Millisecond), + "yGrid", ys) + } else { + slog.Debug("H.264: dropping black warm-up SW frame", + "Y", int(sy), "U", int(su), "V", int(sv), + "streakCount", d.swConsecDroppedFrames, + "streakElapsed", streakElapsed.Round(time.Millisecond)) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + d.swConsecDroppedFrames = 0 + d.swNeedsZeroCheck = false + } + } + if haveRegions { + C.grdp_yuv420p_to_bgra_regions(srcFrame, + (*C.uint8_t)(unsafe.Pointer(&out[0])), C.int(w*4), fullRange, + (*C.uint16_t)(unsafe.Pointer(®ionHint[0])), nRegions) + } else { + dstPtr := (*C.uint8_t)(unsafe.Pointer(&out[0])) + parallelConvertRows(int(h), func(s, e C.int) { + C.grdp_yuv420p_to_bgra_rows(srcFrame, dstPtr, C.int(w*4), fullRange, s, e) + }) + } + + case srcFmt == C.AV_PIX_FMT_NV12 && (!useSwscale || haveRegions): + fullRange := C.int(0) + if srcFrame.color_range == 2 { // AVCOL_RANGE_JPEG + fullRange = 1 + } + frameIdx := d.hwFrameCount + logThis := d.hwFrameCount < 3 + needZeroCheck := d.useHW && d.hwNeedsZeroCheck + var sy, su, sv C.uint8_t + if d.hwFrameCount < 3 || needZeroCheck { + C.grdp_sample_nv12(srcFrame, &sy, &su, &sv) + if d.hwFrameCount < 3 { + slog.Debug("H.264: frame sample (nv12)", + "hw", d.useHW, + "frame", d.hwFrameCount, + "fmt", int(srcFmt), + "colorRange", int(srcFrame.color_range), + "fullRange", int(fullRange), + "Y", int(sy), "U", int(su), "V", int(sv), + "w", int(w), "h", int(h)) + d.hwFrameCount++ + } + } + // Zero-filled IOSurface detection: VideoToolbox sometimes returns an + // uninitialised (all-zero) IOSurface on the first decoded frame after + // decoder init or avcodec_flush_buffers. BT.601 limited-range + // conversion of (Y=0, U=0, V=0) yields BGRA(0,135,0,255), a + // full-screen dark-green frame. Valid NV12 chroma always centres on + // 128, so U=0 and V=0 simultaneously at the centre pixel is an + // unambiguous indicator of an uninitialised buffer. Drop the frame + // and keep hwNeedsZeroCheck set so we continue checking until a frame + // with valid chroma arrives. + // + // A second pattern is Y=0, U≈128, V≈128: VideoToolbox occasionally + // outputs a fully-black frame (Y plane zeroed, UV neutral) for the + // first 1-2 frames after an IDR flush while the pipeline warms up. + // This is detected by sampling Y at a 3×3 grid; if all 9 samples are + // near zero the frame is an uninitialised buffer and is dropped. + if needZeroCheck { + drop := false + if su == 0 && sv == 0 { + drop = true + slog.Debug("H.264: dropping zero-UV HW frame (IOSurface not ready)", + "Y", int(sy)) + } else if sy == 0 && su >= 124 && su <= 132 && sv >= 124 && sv <= 132 && + C.grdp_is_warmup_nv12(srcFrame, 4) != 0 { + drop = true + slog.Debug("H.264: dropping black warm-up HW frame", + "Y", int(sy), "U", int(su), "V", int(sv)) + } + if drop { + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + // Return a non-nil Dropped frame so Decode() counts this as a + // successful VideoToolbox output (health tracking stays correct) + // and callers skip the keyframe-request / decoder-broken path. + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + // Valid frame seen — IOSurface is properly populated. + d.hwNeedsZeroCheck = false + } + // Permanent low-chroma guard for the BGRA NV12 path. The single-pixel + // zero-UV check above is only armed around init/flush/IDR; stale IDR + // priming and other decoder corruption can produce green-monochrome + // frames later in the session. + if C.grdp_is_low_chroma_nv12(srcFrame, lowChromaThreshold, 6) != 0 { + slog.Debug("H.264: dropping low-chroma NV12 frame in BGRA path (green-monochrome corruption)") + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + return &rdpgfx.H264Frame{Dropped: true, Width: int(w), Height: int(h)}, transferNs, nil + } + if haveRegions { + C.grdp_nv12_to_bgra_regions(srcFrame, + (*C.uint8_t)(unsafe.Pointer(&out[0])), C.int(w*4), fullRange, + (*C.uint16_t)(unsafe.Pointer(®ionHint[0])), nRegions) + } else { + dstPtr := (*C.uint8_t)(unsafe.Pointer(&out[0])) + parallelConvertRows(int(h), func(s, e C.int) { + C.grdp_nv12_to_bgra_rows(srcFrame, dstPtr, C.int(w*4), fullRange, s, e) + }) + } + if logThis { + // Sample NV12 input and BGRA output at multiple positions for + // the first three frames to diagnose colour conversion. + for _, p := range [][2]int{{100, 50}, {500, 50}, {960, 50}, {1400, 50}, {960, 200}} { + px, py := p[0], p[1] + if px >= int(w) || py >= int(h) { + continue + } + var sy, su, sv C.uint8_t + C.grdp_sample_nv12_at(srcFrame, C.int(px), C.int(py), &sy, &su, &sv) + off := (py*int(w) + px) * 4 + slog.Debug("H.264: pixel sample (nv12→bgra)", + "frame", frameIdx, + "hw", d.useHW, + "x", px, "y", py, + "Y", int(sy), "U", int(su), "V", int(sv), + "B", out[off], "G", out[off+1], "R", out[off+2]) + } + } + + default: + // For other formats, use swscale. + swsFmt := C.grdp_yuvj_to_yuv(srcFmt) + fullRange := C.grdp_is_full_range_fmt(srcFmt) + if fullRange == 0 && srcFrame.color_range == 2 { // AVCOL_RANGE_JPEG + fullRange = 1 + } + if d.hwFrameCount < 3 { + slog.Debug("H.264: frame sample (swscale)", + "hw", d.useHW, + "frame", d.hwFrameCount, + "fmt", int(srcFmt), + "colorRange", int(srcFrame.color_range), + "fullRange", int(fullRange), + "w", int(w), "h", int(h)) + d.hwFrameCount++ + } + if w != d.lastW || h != d.lastH || srcFmt != d.lastFmt || fullRange != d.lastFullRange { + if d.swsCtx != nil { + C.sws_freeContext(d.swsCtx) + } + d.swsCtx = C.sws_getContext( + w, h, swsFmt, + w, h, C.AV_PIX_FMT_BGRA, + C.SWS_FAST_BILINEAR, nil, nil, nil, + ) + if d.swsCtx == nil { + convErr = fmt.Errorf("sws_getContext failed for %dx%d fmt=%d", w, h, srcFmt) + break + } + C.grdp_sws_set_src_range(d.swsCtx, fullRange) + d.lastW = w + d.lastH = h + d.lastFmt = srcFmt + d.lastFullRange = fullRange + } + C.grdp_frame_to_bgra(d.swsCtx, srcFrame, + (*C.uint8_t)(unsafe.Pointer(&out[0])), C.int(w*4)) + } + + // Extract I420 when explicitly requested (DecodeWithI420) or when in NV12 + // mode but the source is not NV12 (e.g. software decoder producing YUV420P). + // The latter lets callers use LastI420() to refresh the AVC444 Y cache even + // when DecodeWithNV12 returns a BGRA frame instead of native NV12 planes. + if convErr == nil && (d.outI420Enabled || d.outNV12Enabled) { + d.extractI420fromSrc(srcFrame) + } + if usedMapFrame { + C.av_frame_unref(d.mapFrame) + } else if srcFrame == d.swFrame { + C.av_frame_unref(d.swFrame) + } + if convErr != nil { + return nil, transferNs, convErr + } + return &rdpgfx.H264Frame{Data: out, Width: int(w), Height: int(h)}, transferNs, nil +} + +func (d *ffmpegDecoder) Close() { + // Stop any background timers so their callbacks don't fire after Close. + d.stopTimers() + if d.swsCtx != nil { + C.sws_freeContext(d.swsCtx) + d.swsCtx = nil + } + if d.frame != nil { + C.av_frame_free(&d.frame) + } + if d.swFrame != nil { + C.av_frame_free(&d.swFrame) + } + if d.mapFrame != nil { + C.av_frame_free(&d.mapFrame) + } + if d.packet != nil { + C.av_packet_free(&d.packet) + } + if d.codecCtx != nil { + C.avcodec_free_context(&d.codecCtx) + } +} + +func init() { + rdpgfx.SetH264Backend(&rdpgfx.H264DecoderBackend{ + NewHW: func(ch chan<- struct{}) rdpgfx.H264Decoder { + return newH264DecoderInternal(ch, false, keyframeWaitTimeout) + }, + NewSW: func() rdpgfx.H264Decoder { + return newH264DecoderInternal(nil, true, keyframeWaitTimeoutSW) + }, + NewSWFallback: func(ch chan<- struct{}) rdpgfx.H264Decoder { + return newH264DecoderInternal(ch, true, keyframeWaitTimeoutSWFallback) + }, + }) +} diff --git a/plugin/rdpgfx/h264_decoder.go b/plugin/rdpgfx/h264_decoder.go new file mode 100644 index 0000000..7ac643a --- /dev/null +++ b/plugin/rdpgfx/h264_decoder.go @@ -0,0 +1,156 @@ +package rdpgfx + +import "time" + +// H264BrokenReason describes why a decoder became unrecoverable. +type H264BrokenReason int + +const ( + H264BrokenReasonNone H264BrokenReason = iota + H264BrokenReasonInitFailure + H264BrokenReasonHWStall + H264BrokenReasonNoIDR +) + +func (r H264BrokenReason) String() string { + switch r { + case H264BrokenReasonInitFailure: + return "init-failure" + case H264BrokenReasonHWStall: + return "hw-stall" + case H264BrokenReasonNoIDR: + return "no-idr" + default: + return "none" + } +} + +// H264Frame holds a decoded H.264 frame in BGRA pixel format. +type H264Frame struct { + Data []byte // BGRA pixel data, 4 bytes per pixel (nil when Dropped) + Width, Height int + // Dropped is true when the decoder intentionally discarded this frame + // (e.g. zero-filled VideoToolbox IOSurface) rather than experiencing a + // genuine codec stall. Callers must skip bitmap updates but must NOT + // request a keyframe or flag the decoder as broken. + Dropped bool +} + +// H264FrameI420 holds a decoded H.264 frame in planar I420 (YUV420P) format. +// SDL2 can render I420 natively via hardware-accelerated YUV→RGB shaders using +// a PIXELFORMAT_IYUV texture, eliminating CPU-side colour conversion. +// Plane slices borrow ring-buffer memory; the caller must copy all slices +// before the next Decode call. +type H264FrameI420 struct { + Y, U, V []byte + YStride, UStride, VStride int + Width, Height int + FullRange bool // true when the source used full-range (JPEG/PC) YUV +} + +// H264FrameNV12 holds a decoded H.264 frame in NV12 format (Y plane plus +// interleaved UV plane). VideoToolbox commonly transfers hardware-decoded +// H.264 frames as NV12; SDL2 can upload NV12 directly, avoiding the CPU-side +// NV12->I420 deinterleave needed by the I420 path. +type H264FrameNV12 struct { + Y, UV []byte + YStride int + UVStride int + Width int + Height int + FullRange bool +} + +// I420Decoder is an optional interface that an H264Decoder may implement to +// produce I420 output alongside the normal BGRA frame. Callers detect support +// via a type assertion. +type I420Decoder interface { + // DecodeWithI420 decodes H.264 NAL data and returns both a BGRA frame + // and an optional I420 frame for GPU-accelerated rendering. The I420 + // frame is nil when the pixel format is not directly convertible; + // callers must fall back to the BGRA frame in that case. + DecodeWithI420(h264Data []byte) (*H264Frame, *H264FrameI420, error) +} + +// NV12Decoder is an optional interface for decoders that can expose native +// NV12 output. Callers should fall back to I420 or BGRA when the returned +// NV12 frame is nil. +type NV12Decoder interface { + DecodeWithNV12(h264Data []byte) (*H264Frame, *H264FrameNV12, error) +} + +// RegionHinter is an optional interface implemented by decoders that support +// region-aware YUV→BGRA conversion. When SetRegionHint is called immediately +// before Decode, the decoder only converts pixels within the specified dirty +// rectangles, skipping unchanged areas of the frame. +// Each element of rects is [left, top, right, bottom]. +type RegionHinter interface { + SetRegionHint(rects [][4]uint16) +} + +// H264Decoder decodes H.264 Annex B bitstream data into BGRA frames. +type H264Decoder interface { + // Decode decodes H.264 NAL units and returns a decoded frame. + // Returns nil frame (no error) when the decoder needs more input data. + Decode(h264Data []byte) (*H264Frame, error) + // NeedsKeyframe reports whether the decoder is waiting for a keyframe. + NeedsKeyframe() bool + // NeedsIDR reports whether the decoder is explicitly waiting for an IDR frame. + NeedsIDR() bool + // IsBroken reports whether the decoder is permanently unrecoverable. + IsBroken() bool + // BrokenReason reports why the decoder became unrecoverable. + BrokenReason() H264BrokenReason + // ForceBroken marks the decoder unrecoverable for the given reason. + ForceBroken(reason H264BrokenReason) + // HardResetCount returns the number of hard resets performed so far. + HardResetCount() int + // LastReceiveTime returns the wall-clock time of the most recent Decode() call. + LastReceiveTime() time.Time + // Close releases all resources held by the decoder. + Close() +} + +// H264DecoderBackend holds factory functions for creating H264Decoder instances. +// Register a backend via SetH264Backend before starting any RDP session. +// Typically called from an init() function in the application binary. +type H264DecoderBackend struct { + // NewHW creates a hardware-preferred decoder with an optional watchdog channel. + // watchdogCh may be nil for the initial decoder (no watchdog). + NewHW func(watchdogCh chan<- struct{}) H264Decoder + // NewSW creates a software-only decoder without a watchdog (for aux decoders). + NewSW func() H264Decoder + // NewSWFallback creates a software-only decoder with a watchdog + // (used as the post-VideoToolbox-stall fallback decoder). + NewSWFallback func(watchdogCh chan<- struct{}) H264Decoder +} + +var h264Backend *H264DecoderBackend + +// SetH264Backend registers the H.264 decoder backend. +// Must be called before any RDP session is started. +// When the h264 build tag is set, example/h264_ffmpeg.go calls this in its init(). +func SetH264Backend(b *H264DecoderBackend) { h264Backend = b } + +func newH264Decoder() H264Decoder { return newH264DecoderWithWatchdog(nil) } + +func newH264DecoderWithWatchdog(ch chan<- struct{}) H264Decoder { + if h264Backend == nil || h264Backend.NewHW == nil { + return nil + } + return h264Backend.NewHW(ch) +} + +func newH264DecoderSW() H264Decoder { + if h264Backend == nil || h264Backend.NewSW == nil { + return nil + } + return h264Backend.NewSW() +} + +func newH264DecoderSWWithWatchdog(ch chan<- struct{}) H264Decoder { + if h264Backend == nil || h264Backend.NewSWFallback == nil { + return nil + } + return h264Backend.NewSWFallback(ch) +} diff --git a/plugin/rdpgfx/h264_scan.go b/plugin/rdpgfx/h264_scan.go new file mode 100644 index 0000000..331292d --- /dev/null +++ b/plugin/rdpgfx/h264_scan.go @@ -0,0 +1,67 @@ +package rdpgfx + +import "bytes" + +// ScanResult holds the IDR-presence flag and SPS/PPS NAL boundaries (offsets +// into the original packet, including Annex B start code) discovered during a +// single linear walk of an Annex B H.264 packet. +type ScanResult struct { + HasKeyFrame bool + SPSStart, SPSEnd int + PPSStart, PPSEnd int +} + +// ScanH264Packet walks an Annex B H.264 packet exactly once, returning +// whether it contains any IDR slice (NAL type 5) or SPS (NAL type 7) NAL +// unit and recording the byte ranges for the most recent SPS/PPS NALs found. +func ScanH264Packet(data []byte) ScanResult { + var r ScanResult + startCode := []byte{0, 0, 1} + pos := 0 + for pos < len(data) { + off := bytes.Index(data[pos:], startCode) + if off < 0 { + break + } + i := pos + off + scLen := 3 + if i > 0 && data[i-1] == 0 { + i-- + scLen = 4 + } + if i+scLen >= len(data) { + break + } + nalType := data[i+scLen] & 0x1F + if nalType == 5 || nalType == 7 { + r.HasKeyFrame = true + } + if nalType == 7 || nalType == 8 { + searchFrom := i + scLen + 1 + j := len(data) + if searchFrom < len(data) { + if next := bytes.Index(data[searchFrom:], startCode); next >= 0 { + j = searchFrom + next + if j > 0 && data[j-1] == 0 { + j-- + } + } + } + if nalType == 7 { + r.SPSStart, r.SPSEnd = i, j + } else { + r.PPSStart, r.PPSEnd = i, j + } + pos = j + continue + } + pos = i + scLen + 1 + } + return r +} + +// h264PacketHasIDR reports whether an Annex B H.264 packet contains an IDR +// (keyframe) NAL unit. +func h264PacketHasIDR(data []byte) bool { + return ScanH264Packet(data).HasKeyFrame +} diff --git a/plugin/rdpgfx/ict_arm64.go b/plugin/rdpgfx/ict_arm64.go new file mode 100644 index 0000000..70dd731 --- /dev/null +++ b/plugin/rdpgfx/ict_arm64.go @@ -0,0 +1,37 @@ +//go:build arm64 + +package rdpgfx + +import "unsafe" + +// ictToBGRANEON processes n pixels (n must be a multiple of 8) converting +// int16 YCbCr planes to packed BGRA using the ICT formula. Implemented as +// ARMv8-A NEON assembly for throughput. +// +//go:noescape +func ictToBGRANEON(y, cb, cr, dst unsafe.Pointer, n int) + +// ictToBGRA dispatches to the NEON path for full 8-pixel batches and falls +// back to scalar arithmetic for any remaining pixels. +func ictToBGRA(yRow, cbRow, crRow []int16, dst []byte, n int) { + full := (n / 8) * 8 + if full > 0 { + ictToBGRANEON( + unsafe.Pointer(&yRow[0]), + unsafe.Pointer(&cbRow[0]), + unsafe.Pointer(&crRow[0]), + unsafe.Pointer(&dst[0]), + full, + ) + } + for col := full; col < n; col++ { + yv := int32(yRow[col]) + cb := int32(cbRow[col]) + cr := int32(crRow[col]) + ys := (yv + 4096) << 16 + bv := uint32(max(0, min((cb*115992+ys)>>21, 255))) + gv := uint32(max(0, min((ys-cb*22527-cr*46819)>>21, 255))) + rv := uint32(max(0, min((cr*91916+ys)>>21, 255))) + *(*uint32)(unsafe.Pointer(&dst[col*4])) = bv | gv<<8 | rv<<16 | 0xFF000000 + } +} diff --git a/plugin/rdpgfx/ict_arm64.s b/plugin/rdpgfx/ict_arm64.s new file mode 100644 index 0000000..91df9f8 --- /dev/null +++ b/plugin/rdpgfx/ict_arm64.s @@ -0,0 +1,153 @@ +// ARM64 NEON implementation of the ICT (Irreversible Color Transform) inverse +// for RemoteFX tiles: converts Y/Cb/Cr int16 planes to packed BGRA uint8. +// +// The Go arm64 assembler does not expose SSHLL, SSHR, SQXTN, SQXTUN, MUL, or +// MLA as named mnemonics, so each of those instructions is emitted as a raw +// 32-bit WORD constant with the ARMv8-A encoding. All encodings were derived +// by assembling the equivalent C-style mnemonics with the system `as` tool and +// confirmed by inspecting the resulting object file. +// +// Register map: +// R0 – y plane pointer (advances 16 bytes per iteration) +// R1 – cb plane pointer +// R2 – cr plane pointer +// R3 – dst BGRA pointer (advances 32 bytes per iteration) +// R4 – iteration count (= n/8 on entry) +// R5 – scratch GP register (used only during constant setup) +// +// V0 – y[0..7] int16 raw (8H) +// V1 – cb[0..7] int16 raw (8H) +// V2 – cr[0..7] int16 raw (8H) +// V3 – y_lo[0..3] int32 (4S) | sign-extended from lower 4 lanes of V0 +// V4 – y_hi[4..7] int32 (4S) | sign-extended from upper 4 lanes of V0 +// V5 – cb_lo int32 (4S) +// V6 – cb_hi int32 (4S) +// V7 – cr_lo int32 (4S) +// V8 – cr_hi int32 (4S) +// V9 – ys_lo = (y_lo + 4096) << 16 (4S) – intermediate +// V10 – ys_hi (4S) +// V11 – const { 4096, ... } (4S) ─┐ +// V12 – const { 115992, ... } │ loaded once before loop +// V13 – const { 22527, ... } │ +// V14 – const { 46819, ... } │ +// V15 – const { 91916, ... } (4S) ┘ +// V16 – B lo (4S), then packed B (4H → 8B) +// V17 – B hi (4S) +// V18 – G lo (4S) +// V19 – G hi (4S) +// V20 – R lo (4S) +// V21 – R hi (4S) +// V25 – B uint8[8] after saturation +// V26 – G uint8[8] +// V27 – R uint8[8] +// V28 – alpha = 0xFF (8B) – constant, preserved across iterations +// +// Stack ABI (ABI0): +// y+0(FP) unsafe.Pointer +// cb+8(FP) unsafe.Pointer +// cr+16(FP) unsafe.Pointer +// dst+24(FP) unsafe.Pointer +// n+32(FP) int + +#include "textflag.h" + +// func ictToBGRANEON(y, cb, cr, dst unsafe.Pointer, n int) +TEXT ·ictToBGRANEON(SB),NOSPLIT,$0-40 + MOVD y+0(FP), R0 + MOVD cb+8(FP), R1 + MOVD cr+16(FP), R2 + MOVD dst+24(FP), R3 + MOVD n+32(FP), R4 + + // alpha constant: V28.8B = {0xFF, ...} + WORD $0x0f07e7fc // movi v28.8b, #0xff + + // Load ICT coefficients into V11-V15. + MOVD $4096, R5 + WORD $0x4e040cab // dup v11.4s, w5 (Y bias) + MOVD $115992, R5 + WORD $0x4e040cac // dup v12.4s, w5 (Cb → B) + MOVD $22527, R5 + WORD $0x4e040cad // dup v13.4s, w5 (Cb → G, subtracted) + MOVD $46819, R5 + WORD $0x4e040cae // dup v14.4s, w5 (Cr → G, subtracted) + MOVD $91916, R5 + WORD $0x4e040caf // dup v15.4s, w5 (Cr → R) + + // R4 = number of 8-pixel batches. + LSR $3, R4, R4 + CBZ R4, done + +loop: + // ── Load 8 int16 pixels from each plane ────────────────────────────── + WORD $0x4cdf7400 // ld1 {v0.8h}, [x0], #16 + WORD $0x4cdf7421 // ld1 {v1.8h}, [x1], #16 + WORD $0x4cdf7442 // ld1 {v2.8h}, [x2], #16 + + // ── Sign-extend int16 → int32 ───────────────────────────────────────── + WORD $0x0f10a403 // sshll v3.4s, v0.4h, #0 y_lo + WORD $0x4f10a404 // sshll2 v4.4s, v0.8h, #0 y_hi + WORD $0x0f10a425 // sshll v5.4s, v1.4h, #0 cb_lo + WORD $0x4f10a426 // sshll2 v6.4s, v1.8h, #0 cb_hi + WORD $0x0f10a447 // sshll v7.4s, v2.4h, #0 cr_lo + WORD $0x4f10a448 // sshll2 v8.4s, v2.8h, #0 cr_hi + + // ── ys = (y + 4096) << 16 ───────────────────────────────────────────── + WORD $0x4eab8469 // add v9.4s, v3.4s, v11.4s + WORD $0x4eab848a // add v10.4s, v4.4s, v11.4s + WORD $0x4f305529 // shl v9.4s, v9.4s, #16 + WORD $0x4f30554a // shl v10.4s, v10.4s, #16 + + // ── B = ys + cb * 115992 ────────────────────────────────────────────── + WORD $0x4ea91d30 // mov v16.16b, v9.16b (b_lo = ys_lo) + WORD $0x4eaa1d51 // mov v17.16b, v10.16b (b_hi = ys_hi) + WORD $0x4eac94b0 // mla v16.4s, v5.4s, v12.4s + WORD $0x4eac94d1 // mla v17.4s, v6.4s, v12.4s + + // ── G = ys - cb*22527 - cr*46819 ────────────────────────────────────── + WORD $0x4ea91d32 // mov v18.16b, v9.16b (g_lo = ys_lo) + WORD $0x4eaa1d53 // mov v19.16b, v10.16b (g_hi = ys_hi) + WORD $0x4ead9ca0 // mul v0.4s, v5.4s, v13.4s (cb_lo*22527) + WORD $0x6ea08652 // sub v18.4s, v18.4s, v0.4s + WORD $0x4eae9ce0 // mul v0.4s, v7.4s, v14.4s (cr_lo*46819) + WORD $0x6ea08652 // sub v18.4s, v18.4s, v0.4s + WORD $0x4ead9cc1 // mul v1.4s, v6.4s, v13.4s (cb_hi*22527) + WORD $0x6ea18673 // sub v19.4s, v19.4s, v1.4s + WORD $0x4eae9d01 // mul v1.4s, v8.4s, v14.4s (cr_hi*46819) + WORD $0x6ea18673 // sub v19.4s, v19.4s, v1.4s + + // ── R = ys + cr * 91916 ─────────────────────────────────────────────── + WORD $0x4ea91d34 // mov v20.16b, v9.16b (r_lo = ys_lo) + WORD $0x4eaa1d55 // mov v21.16b, v10.16b (r_hi = ys_hi) + WORD $0x4eaf94f4 // mla v20.4s, v7.4s, v15.4s + WORD $0x4eaf9515 // mla v21.4s, v8.4s, v15.4s + + // ── Arithmetic shift right by 21 ────────────────────────────────────── + WORD $0x4f2b0610 // sshr v16.4s, v16.4s, #21 + WORD $0x4f2b0631 // sshr v17.4s, v17.4s, #21 + WORD $0x4f2b0652 // sshr v18.4s, v18.4s, #21 + WORD $0x4f2b0673 // sshr v19.4s, v19.4s, #21 + WORD $0x4f2b0694 // sshr v20.4s, v20.4s, #21 + WORD $0x4f2b06b5 // sshr v21.4s, v21.4s, #21 + + // ── Narrow int32 → int16 with signed saturation ─────────────────────── + WORD $0x0e614a19 // sqxtn v25.4h, v16.4s B lo + WORD $0x4e614a39 // sqxtn2 v25.8h, v17.4s B hi + WORD $0x0e614a5a // sqxtn v26.4h, v18.4s G lo + WORD $0x4e614a7a // sqxtn2 v26.8h, v19.4s G hi + WORD $0x0e614a9b // sqxtn v27.4h, v20.4s R lo + WORD $0x4e614abb // sqxtn2 v27.8h, v21.4s R hi + + // ── Narrow int16 → uint8, clamping to [0, 255] ──────────────────────── + WORD $0x2e212b39 // sqxtun v25.8b, v25.8h B + WORD $0x2e212b5a // sqxtun v26.8b, v26.8h G + WORD $0x2e212b7b // sqxtun v27.8b, v27.8h R + + // ── Store 8 BGRA pixels interleaved, advance dst by 32 ──────────────── + WORD $0x0c9f0079 // st4 {v25.8b, v26.8b, v27.8b, v28.8b}, [x3], #32 + + SUBS $1, R4, R4 + BNE loop + +done: + RET diff --git a/plugin/rdpgfx/ict_generic.go b/plugin/rdpgfx/ict_generic.go new file mode 100644 index 0000000..aad3ec9 --- /dev/null +++ b/plugin/rdpgfx/ict_generic.go @@ -0,0 +1,38 @@ +//go:build !arm64 + +package rdpgfx + +import "unsafe" + +// ictToBGRA converts n pixels from YCbCr (ICT) to BGRA and writes them into +// dst (which must hold ≥ 4*n bytes). Processing n pixels in [8]int32 arrays +// with a scalar inner loop. +func ictToBGRA(yRow, cbRow, crRow []int16, dst []byte, n int) { + const batch = 8 + full := (n / batch) * batch + for base := 0; base < full; base += batch { + var yv, cb, cr [batch]int32 + for k := range batch { + yv[k] = int32(yRow[base+k]) + cb[k] = int32(cbRow[base+k]) + cr[k] = int32(crRow[base+k]) + } + for k := range batch { + ys := (yv[k] + 4096) << 16 + bv := uint32(max(0, min((cb[k]*115992+ys)>>21, 255))) + gv := uint32(max(0, min((ys-cb[k]*22527-cr[k]*46819)>>21, 255))) + rv := uint32(max(0, min((cr[k]*91916+ys)>>21, 255))) + *(*uint32)(unsafe.Pointer(&dst[(base+k)*4])) = bv | gv<<8 | rv<<16 | 0xFF000000 + } + } + for col := full; col < n; col++ { + yv := int32(yRow[col]) + cb := int32(cbRow[col]) + cr := int32(crRow[col]) + ys := (yv + 4096) << 16 + bv := uint32(max(0, min((cb*115992+ys)>>21, 255))) + gv := uint32(max(0, min((ys-cb*22527-cr*46819)>>21, 255))) + rv := uint32(max(0, min((cr*91916+ys)>>21, 255))) + *(*uint32)(unsafe.Pointer(&dst[col*4])) = bv | gv<<8 | rv<<16 | 0xFF000000 + } +} diff --git a/plugin/rdpgfx/rdpgfx.go b/plugin/rdpgfx/rdpgfx.go new file mode 100644 index 0000000..c98d783 --- /dev/null +++ b/plugin/rdpgfx/rdpgfx.go @@ -0,0 +1,2816 @@ +package rdpgfx + +import ( + "encoding/binary" + "fmt" + "log/slog" + "runtime/debug" + "sync" + "sync/atomic" + "time" + + "git.zeroonesoft.cn/golib/rdplib/plugin" +) + +// regionPool reuses byte slices for progressive codec rectangle extraction, +// avoiding per-rectangle allocations that cause GC pressure. +var regionPool = sync.Pool{ + New: func() any { return []byte(nil) }, +} + +const ( + ChannelName = plugin.RDPGFX_DVC_CHANNEL_NAME +) + +// suspendFrameAcknowledge is the queueDepth value defined in MS-RDPEGFX +// 2.2.2.8 that instructs the server to suspend sending new frames until +// the client sends a subsequent FRAME_ACKNOWLEDGE with a different value. +const suspendFrameAcknowledge uint32 = 0xFFFFFFFF + +// RDPGFX Command IDs (MS-RDPEGFX 2.2.2) +const ( + cmdidWireToSurface1 uint16 = 0x0001 + cmdidWireToSurface2 uint16 = 0x0002 + cmdidDeleteEncodingContext uint16 = 0x0003 + cmdidSolidFill uint16 = 0x0004 + cmdidSurfaceToSurface uint16 = 0x0005 + cmdidSurfaceToCache uint16 = 0x0006 + cmdidCacheToSurface uint16 = 0x0007 + cmdidEvictCacheEntry uint16 = 0x0008 + cmdidCreateSurface uint16 = 0x0009 + cmdidDeleteSurface uint16 = 0x000A + cmdidStartFrame uint16 = 0x000B + cmdidEndFrame uint16 = 0x000C + cmdidFrameAcknowledge uint16 = 0x000D + cmdidResetGraphics uint16 = 0x000E + cmdidMapSurfaceToOutput uint16 = 0x000F + cmdidCacheImportOffer uint16 = 0x0010 + cmdidCacheImportReply uint16 = 0x0011 + cmdidCapsAdvertise uint16 = 0x0012 + cmdidCapsConfirm uint16 = 0x0013 + cmdidMapSurfaceToWindow uint16 = 0x0015 + cmdidQoeFrameAcknowledge uint16 = 0x0016 + cmdidMapSurfaceToScaledOutput uint16 = 0x0017 + cmdidMapSurfaceToScaledWindow uint16 = 0x0018 +) + +// Pixel Formats +const ( + pixelFormatXRGB8888 uint8 = 0x20 + pixelFormatARGB8888 uint8 = 0x21 +) + +// Codec IDs (MS-RDPEGFX 2.2.2.1 / FreeRDP rdpgfx.h) +const ( + codecUncompressed uint16 = 0x0000 + codecCaVideo uint16 = 0x0003 // RDPGFX_CODECID_CAVIDEO (RemoteFX tiles) + codecPlanar uint16 = 0x0004 + codecClear uint16 = 0x0008 + codecProgressive uint16 = 0x0009 + codecAVC420 uint16 = 0x000B + codecAVC444 uint16 = 0x000E + codecAVC444v2 uint16 = 0x000F +) + +// Capability versions and flags +const ( + capVersion8 uint32 = 0x00080004 + capVersion81 uint32 = 0x00080105 + capVersion10 uint32 = 0x000A0002 + capVersion101 uint32 = 0x000A0100 + capVersion102 uint32 = 0x000A0200 + capVersion103 uint32 = 0x000A0301 + capVersion104 uint32 = 0x000A0400 + capVersion105 uint32 = 0x000A0502 + capVersion106 uint32 = 0x000A0600 + capVersion1061 uint32 = 0x000A0601 + capVersion107 uint32 = 0x000A0701 + capFlagThinClient uint32 = 0x00000001 + capFlagSmallCache uint32 = 0x00000002 + capFlagAVC420Enabled uint32 = 0x00000010 // v8.1: explicitly enable AVC420 + capFlagAVCDisabled uint32 = 0x00000020 // v10+: disable AVC +) + +const headerSize = 8 + +// BitmapUpdate represents a rendered bitmap region. +// +// Lifecycle: Data is borrowed from an internal buffer pool and is only +// valid for the duration of the synchronous onBitmap callback. After the +// callback returns, the slice may be returned to the pool and overwritten +// by subsequent updates. Callers that need to retain the pixels (e.g. to +// hand them to an asynchronous paint goroutine) MUST copy the bytes +// before the callback returns. +type BitmapUpdate struct { + DestLeft, DestTop, DestRight, DestBottom int + Width, Height int + Bpp int // bytes per pixel (always 4) + Data []byte // BGRA pixel data — see lifecycle note above +} + +// bitmapBufPool reuses BGRA byte slices used to back BitmapUpdate.Data. +// Buffers are acquired with acquireBitmapBuf, handed to the onBitmap +// callback, and released with releaseBitmapBuf once the (synchronous) +// callback returns. This eliminates per-rectangle allocations on the +// hot CaVideo / AVC partial-blit paths. +var bitmapBufPool = sync.Pool{ + New: func() any { return []byte(nil) }, +} + +// decodePkt is the message type for the async decode channel. +// pooled is true when data was acquired from bitmapBufPool; the receiver +// must call releaseBitmapBuf(data) after processing. +type decodePkt struct { + data []byte + pooled bool +} + +func acquireBitmapBuf(size int) []byte { + if size <= 0 { + return nil + } + b := bitmapBufPool.Get().([]byte) + if cap(b) < size { + return make([]byte, size) + } + return b[:size] +} + +func releaseBitmapBuf(b []byte) { + if b == nil { + return + } + //nolint:staticcheck // intentional pool of byte slices + bitmapBufPool.Put(b[:cap(b)]) +} + +// emitAndReleaseUpdates calls the onBitmap callback and then returns the +// pooled Data buffers of the supplied updates back to bitmapBufPool. All +// updates passed in must have Data acquired via acquireBitmapBuf. +func (g *GfxHandler) emitAndReleaseUpdates(updates []BitmapUpdate) { + if g.onBitmap != nil && len(updates) > 0 { + g.onBitmap(updates) + } + for i := range updates { + releaseBitmapBuf(updates[i].Data) + updates[i].Data = nil + } +} + +type surface struct { + width, height uint16 + format uint8 + data []byte // BGRA, 4 bytes per pixel + outputX uint32 + outputY uint32 + mapped bool + // shadowStale is true when the CPU surface shadow (data) may not match the + // pixels currently shown on the GPU display. It is set when an AVC frame + // advances the GPU display (onNV12/onI420) without a corresponding BGRA + // shadow update (decoded == nil), and cleared by a full-surface blit. While + // stale, the region-only shadow blit fast path is bypassed in favour of a + // full blit so the shadow is fully repaired before the next partial update. + // Initialised true so the first frame on a fresh (zero-filled) surface always + // takes the full-blit path. + shadowStale bool +} + +type vBarEntry struct { + pixels []byte // BGRA pixel data, 4 bytes per pixel + count int +} + +type cacheEntry struct { + data []byte // BGRA pixel data + width, height int + key uint64 // 服务器提供的持久缓存键(SurfaceToCache),0 = 无 +} + +// GfxCacheEntry 是一条跨连接保留的位图缓存条目(MS-RDPEGFX 持久位图缓存)。 +// Key 为服务器在 SurfaceToCache 中给出的 cacheKey(持久身份),Data 为当时 +// 的 BGRA 像素副本。重连后经 CacheImportOffer 上报,服务器按前缀导入并 +// 重新分配槽位,后续 CacheToSurface 即可直接回贴、无需重传像素。 +type GfxCacheEntry struct { + Key uint64 + Width, Height int + Bpp uint16 + Data []byte +} + +// GfxCacheStore 在浏览器侧持久化缓存条目(v1:页面内存,跨手动重连保留)。 +// Persist 随每条 SurfaceToCache 调用;Export 在每次连接 caps 确认后调用一次, +// 返回待上报的条目(顺序即 CacheImportReply 的前缀导入顺序)。 +// Get/Keys 供 bitmap 管线持久缓存(6.4b M2)按需查取与枚举键。 +type GfxCacheStore interface { + Persist(key uint64, w, h int, bpp uint16, data []byte) + Export() []GfxCacheEntry + Get(key uint64) (GfxCacheEntry, bool) + Keys() []uint64 +} + +// GfxHandler implements the RDPGFX (MS-RDPEGFX) protocol. +type GfxHandler struct { + surfaces map[uint16]*surface + cacheEntries map[uint16]cacheEntry + cacheStore GfxCacheStore // 持久缓存桥(nil = 未启用) + offeredCache []GfxCacheEntry // 本连接 CacheImportOffer 的条目,等待 Reply 前缀映射 + importOfferSent bool // 每连接只上报一次 + clearCtx *clearCodecCtx + zgfx *zgfxContext + rfx *rfxDecoder + progressive *rfxProgressiveDecoder + // codecBytes[i] 累计 codecId=i 的表面位图字节数(WTS1+WTS2),用于带宽诊断 + codecBytes [16]atomic.Int64 + // cmdCounts[i] 累计 cmdId=i 的 PDU 条数(诊断服务端停止发帧用) + cmdCounts [64]atomic.Int64 + h264dec H264Decoder + // h264dec2 is the auxiliary H.264 decoder used for AVC444v2 LC=2 chroma-upgrade + // frames. It decodes stream2, which carries chroma values for positions not + // covered by stream1's 4:2:0 quantiser. The decoded I420 planes are combined + // with the luma and chroma planes cached from the most recent LC=0/1 main-stream + // decode to reconstruct full 4:4:4 YUV before converting to BGRA. + h264dec2 H264Decoder + // avc444YPlane caches the luma (Y) and half-res chroma (U/V) planes from the + // last main-stream AVC444 decode, for use when an LC=2 chroma-upgrade frame arrives. + avc444YPlane avc444YPlane + // avc444IDRYPlane caches the luma and half-res chroma planes from the most + // recently decoded stream1 IDR frame. When a standalone LC=2 packet carries + // a stream2 IDR, it should be combined with the matching stream1 IDR luma + // (not the latest P-frame luma stored in avc444YPlane), so we keep this + // separate snapshot. + avc444IDRYPlane avc444YPlane + // lc2SampleLogged is set after the first LC=2 combine output has been + // sampled for green/pink colour diagnostics. Reset on each stream1 IDR so + // we can observe the combine quality at every GOP boundary. + lc2SampleLogged bool + lc2PFrameSampleLogged bool // logged first P-frame LC=2 combine after IDR + lc0SampleLogged bool // logged first LC=0 IDR frame pixel samples + // framesDecoded is accessed from both read and decode goroutines. + framesDecoded atomic.Uint32 + // 帧间隔与单消息解码耗时的 EMA(微秒),DiagStats 输出用。 + // 写方为 decode goroutine,读方为统计循环,CAS 循环更新。 + frameIntvUs atomic.Int64 + decUs atomic.Int64 + lastFrameAt atomic.Int64 // unix nano,0 = 尚未收到帧 + // 帧解码起点(unix ns):StartFrame 到 EndFrame 的墙钟时长进 QoE 上报 + frameDecodeStart atomic.Int64 + // SUSPEND 状态:队列深时置位(向服务器请求暂停发送),排空后复位 + suspended atomic.Bool + // sessionW/H 为通道初始化时发送 RESET_GRAPHICS 用的会话尺寸 + //(grdp.go 创建 handler 时注入)。 + sessionW, sessionH uint32 + sendFn func(data []byte) + onBitmap func([]BitmapUpdate) + // decodeCh receives decompressed PDU data for asynchronous decode. + decodeCh chan decodePkt + // ackCh is a buffered channel of serialized ACK PDUs. Every + // EndFrame ACK is enqueued here and the writeLoop goroutine sends + // each one to the server. The server tracks outstanding frames + // individually, so skipping ACKs causes it to stop sending. + ackCh chan []byte + // doneCh is closed by Close() to signal decodeLoop and writeLoop to exit. + doneCh chan struct{} + closeOnce sync.Once + // onDecoderBroken is called once when the H.264 decoder becomes permanently + // unrecoverable (all soft resets exhausted). The caller should reconnect + // the RDP session to create a fresh decoder. + onDecoderBroken func() + decoderBrokenNotified bool + // watchdogCh receives signals from background timers inside ffmpegDecoder + // when stall-probe or IDR-wait timeouts expire independently of server + // frame rate. decodeLoop selects on this channel so it calls + // maybeNotifyDecoderBroken even when no server frames are arriving. + watchdogCh chan struct{} + // lastDecodedFrame records when a visible AVC frame was last produced. + // Local-input watchdogs compare against this timestamp so recovery does + // not wait for a subsequent H.264 packet to arrive. + lastDecodedFrame atomic.Int64 + inputWatchdogMu sync.Mutex + inputWatchdog *time.Timer + inputWatchdogNS int64 + // lastLC2RecvTime records when the most recent AVC444 LC=2 frame arrived, + // regardless of whether it could be decoded. Used to detect the + // "server sending LC=2 only, aux decoder absent" deadlock. + lastLC2RecvTime atomic.Int64 + // auxDecoderBrokenTimer fires after auxDecoderBrokenTimeout when h264dec2 + // is nil. When it fires it signals watchdogCh so that decodeLoop can call + // maybeRenegotiateCapabilities and break the LC=2-only deadlock. + auxDecoderBrokenTimer *time.Timer + auxDecoderBrokenTimerMu sync.Mutex + // onKeyframeRequest is called to ask the server to send a fresh IDR + // keyframe. Optional: if nil, the decoder will wait for the next + // server-initiated keyframe. + onKeyframeRequest func() + // lastKeyframeRequest is the wall-clock time of the most recent keyframe + // request sent to the server. Used to rate-limit repeat requests. + lastKeyframeRequest time.Time + // softResetCount tracks how many in-place decoder resets have been + // attempted since the last server-triggered RESET_GRAPHICS. + softResetCount int + // noIDRSoftResetCount tracks soft resets triggered specifically by the + // no-IDR broken reason. This counter is kept separate from softResetCount + // so that a prior HW-stall reset (which increments softResetCount) does not + // consume the no-IDR recovery budget. Reset on RESET_GRAPHICS and on a + // successful frame decode. + noIDRSoftResetCount int + // usingSWFallback is set after a HW stall forces a switch to software + // decoding. Both h264dec and h264dec2 are created SW-only while this + // flag is true, avoiding repeated VideoToolbox stalls that would + // otherwise trigger a full RDP reconnect. + usingSWFallback bool + // swFallbackPrimed is set after a SW fallback decoder has been primed + // with the cached (stale) stream1 IDR. It stays armed for the whole SW + // fallback window — until a genuine fresh IDR resyncs the decoder + // (maybeCacheStream1IDR) or RESET_GRAPHICS — so that consecutive dropped + // frames caused by the stale prime can be detected and escalated to a + // reconnect instead of producing endless green/zero-UV frames. It is NOT + // cleared by a single successful decode: stale-prime corruption only + // appears once the following P-frames diverge from the missing reference. + swFallbackPrimed bool + // swFallbackDroppedCount counts dropped frames since the SW fallback + // decoder was primed. If it exceeds swFallbackDropLimit the decoder is + // declared broken and onDecoderBroken is fired. + swFallbackDroppedCount int + // swFallbackFirstDropTime records when the current run of consecutive + // stale-prime drops started. It bounds how long corruption may persist + // before escalating to a reconnect, giving the ForceRefresh resync IDR a + // chance to heal the picture first. Reset whenever a clean frame decodes. + swFallbackFirstDropTime time.Time + // lc2EverDecoded is set to true after the first successful AVC444 LC=2 + // chroma-upgrade decode. maybeRenegotiateCapabilities uses this to + // distinguish "LC=2 was working and then broke" (reconnect needed) from + // "LC=2 never worked this session" (server may not support stream2 priming; + // gracefully degrade to LC=0 only without reconnecting). + lc2EverDecoded bool + // auxDecoderNoIDRRetries counts how many times maybeRenegotiateCapabilities + // has been called in the "stream2EverSeen but lc2EverDecoded=false" case + // this session. Each attempt sends a ForceRefresh; after + // auxDecoderMaxIDRRetries consecutive attempts the session degrades to LC=0 + // only (no reconnect — the server consistently omits stream2 IDRs). + // Reset on RESET_GRAPHICS, when LC=2 successfully decodes, and whenever + // maybeRenegotiateCapabilities returns early due to no recent LC=2 activity + // (e.g. during a HW-decoder GOP-boundary stall) so a resumed burst gets + // fresh retries rather than inheriting the previous count. + auxDecoderNoIDRRetries int + // lc2PermanentlyDegraded is set when the server has not delivered a stream2 + // IDR despite repeated keyframe requests and lc2EverDecoded is still false. + // Once set, LC=2 frames are silently skipped for the remainder of the session + // without arming the renegotiation timer — avoiding an endless reconnect loop. + // Cleared on RESET_GRAPHICS so a fresh AVC444 sequence gets a clean slate. + // Cleared in primeAuxDecoder on a stream2 IDR to allow late recovery. + lc2PermanentlyDegraded bool + // stream2EverSeen is set when a non-empty stream2 payload is observed inside + // an LC=0 packet. VirtualBox VRDE never includes stream2 in LC=0 packets, + // so stream2EverSeen stays false for the whole session. Windows does include + // stream2 in LC=0 IDRs, so stream2EverSeen becomes true as soon as the first + // LC=0 AVC444 packet arrives. maybeRenegotiateCapabilities uses this to + // distinguish "server never sends stream2" (VirtualBox → permanent degrade) + // from "stream2 seen but aux decoder not yet primed" (Windows, Chrome just + // launched → transient, wait for IDR rather than logging a WARN). + // Reset on RESET_GRAPHICS since the server starts a fresh AVC444 sequence. + stream2EverSeen bool + // lastStream1IDR caches the most recent stream1 H.264 IDR NAL data in + // Annex B format. When VideoToolbox stalls and the decoder falls back to + // software, this data is fed immediately to the new SW decoder so it can + // decode subsequent P-frames without waiting for the server to send a fresh + // IDR via ForceRefresh (which some servers, e.g. VirtualBox VRDE, ignore). + // Cleared on RESET_GRAPHICS to avoid feeding a stale IDR to a new pipeline. + lastStream1IDR []byte + // lastStream1IDRTime is the wall-clock time when the most recent stream1 IDR + // was cached, logged for diagnostics (e.g. how stale the priming IDR was). + // lastStream1IDRFrame is framesDecoded at the same instant, logged for diagnostics. + lastStream1IDRTime time.Time + lastStream1IDRFrame uint32 + // avc444Disabled, when true, limits the CAPS_ADVERTISE to v8.0 and v8.1 + // (AVC420 only). The server will never send AVC444/AVC444v2 frames, which + // avoids the LC=2 colour degradation seen with VirtualBox VRDE. + avc444Disabled bool + // avcDisabled, when true, advertises v10.x with AVC_DISABLED so the server + // keeps the RDPGFX channel on ClearCodec/RFX Progressive (no H.264 decode). + avcDisabled bool + // pduRecord, when non-nil, receives every wire-to-surface bitmap payload + // for offline replay (garbled-screen/bandwidth debugging harness). + pduRecord func(kind byte, codecId uint16, surfW, surfH, x, y, w, h uint32, payload []byte) + // queueDepthHint is a minimum queueDepth to report in FRAME_ACKNOWLEDGE + // PDUs. A higher value makes the server believe the client has a larger + // decode backlog, causing it to slow down or reduce encoding quality. + // 0 means "report the real queue length" (default, no throttling). + // See SetQueueDepthHint. + queueDepthHint atomic.Uint32 + // singleUpdate is a pre-allocated one-element slice reused by emitBitmap + // and emitBitmapPooled to avoid a heap allocation on every BGRA frame. + // Safe: all emitBitmap* calls run on the single decode goroutine. + singleUpdate [1]BitmapUpdate + // updatesBuf is a pre-allocated slice reused by emitCaVideoRects and + // blitAndEmitAVCRegions to avoid make() on every multi-region frame. + // Safe: both functions run on the single decode goroutine and never call + // each other. + updatesBuf []BitmapUpdate + // avcStream1 and avcStream2 are pre-allocated AVC stream structs reused by + // fillAVC444Stream and the decode methods to avoid a heap allocation for + // the stream header (struct + regions slice) on every H.264 frame. + // Safe: all AVC decode calls run on the single decode goroutine. + avcStream1 avc420Stream + avcStream2 avc420Stream + // regionHintBuf is a pre-allocated slice reused by the SetRegionHint call + // sites to avoid a make([][4]uint16, ...) on every H.264 frame that carries + // dirty region metadata. + // Safe: used only on the single decode goroutine. + regionHintBuf [][4]uint16 + // onH264Raw is called with raw H.264 NAL unit data when h264dec is nil + // (e.g. WASM builds without CGo). The caller can forward the data to a + // JavaScript WebCodecs VideoDecoder instead. + // destX, destY are the top-left canvas coordinates. + // regions 为扁平 [l,t,r,b,...](帧内坐标,右下开区间):服务器只保证 + // 区域内像素有效,帧内其余像素未定义,绘制端必须只贴区域;空切片 + // 表示整帧有效。 + onH264Raw func(destX, destY, w, h int, isKey bool, data []byte, regions []int32) + // onI420 is called after a successful H.264 decode when I420 planar data + // is available. The caller can upload the planes to an SDL2 IYUV texture + // for GPU-accelerated YUV→RGB conversion, bypassing the CPU colour path. + // destX, destY are absolute canvas coordinates. + onI420 func(destX, destY, w, h int, y []byte, yStride int, u []byte, uStride int, v []byte, vStride int) + // onNV12 is like onI420 but receives native NV12 output. It is preferred + // by SDL2 clients when available because VideoToolbox commonly outputs + // NV12 and SDL_UpdateNVTexture can upload it directly. + onNV12 func(destX, destY, w, h int, y []byte, yStride int, uv []byte, uvStride int) +} + +// NewGfxHandler creates a new RDPGFX handler. +func NewGfxHandler(onBitmap func([]BitmapUpdate)) *GfxHandler { + g := &GfxHandler{ + surfaces: make(map[uint16]*surface), + cacheEntries: make(map[uint16]cacheEntry), + clearCtx: newClearCodecCtx(), + zgfx: newZgfxContext(), + rfx: newRfxDecoder(), + progressive: newRfxProgressiveDecoder(), + // h264dec2 starts nil; primeAuxDecoder creates it on the first stream2 IDR + // so it is always primed before decoding LC=2 P-frames. + onBitmap: onBitmap, + decodeCh: make(chan decodePkt, 64), + ackCh: make(chan []byte, 512), + doneCh: make(chan struct{}), + watchdogCh: make(chan struct{}, 4), + } + g.h264dec = newH264DecoderWithWatchdog(g.watchdogCh) + go g.decodeLoop() + go g.writeLoop() + return g +} + +// inputStallSilentThreshold is the minimum time without a decoded frame before +// the local-input watchdog considers the HW decoder potentially stalled. +// Kept in the non-build-tagged file so both normal and !h264 builds compile. +// Must be kept in sync with avcHWReadyFreezeThreshold in h264_ffmpeg.go. +const inputStallSilentThreshold = 7 * time.Second + +const localInputRecoveryGrace = 750 * time.Millisecond + +// auxDecoderBrokenTimeout is the maximum time we wait for an LC=0 stream2 IDR +// to arrive and recreate the aux decoder (h264dec2) after it has been torn down. +// If LC=2 frames keep arriving beyond this window without an LC=0 IDR, the RDPGFX +// capabilities are re-advertised to force the server to issue RESET_GRAPHICS and +// restart the video pipeline with a fresh LC=0 IDR for both streams. +const auxDecoderBrokenTimeout = 10 * time.Second + +// auxDecoderMaxIDRRetries is the maximum number of ForceRefresh keyframe +// requests sent while waiting for a stream2 IDR before giving up and +// reconnecting. Each attempt is spaced by auxDecoderBrokenTimeout (10 s). +// Windows servers typically begin a new GOP every 30–60 s, so 5 attempts +// (≤50 s) provides enough coverage for at least one natural IDR boundary. +const auxDecoderMaxIDRRetries = 5 + +// Close shuts down the GfxHandler's background goroutines. +// Safe to call multiple times; subsequent calls are no-ops. +// +// h264dec is intentionally NOT freed here: decodeLoop (goroutine 21) may be +// in the middle of avcodec_send_packet when Close is called from the transport +// goroutine, which would cause a use-after-free SIGSEGV. Instead, decodeLoop +// defers cleanup of h264dec so it always runs after the last Decode call. +func (g *GfxHandler) Close() { + g.closeOnce.Do(func() { + g.stopInputWatchdog() + g.stopAuxDecoderBrokenTimer() + close(g.doneCh) + }) +} + +// NotifyLocalInput tells the graphics pipeline that a real local input event +// was just sent to the server. If the decoder has already been silent longer +// than the HW stall threshold, arm a short watchdog so recovery no longer +// depends on the next H.264 packet arriving. +func (g *GfxHandler) NotifyLocalInput() { + if g.h264dec == nil { + return + } + // Don't arm the watchdog when the decoder is waiting for an IDR after a + // reset: it is already in a known recovery state, and arming the watchdog + // here would force-break the newly reset decoder 750 ms later when the IDR + // hasn't arrived yet, causing rapid cascading soft resets. + if g.h264dec.NeedsIDR() { + return + } + lastDecodedNS := g.lastDecodedFrame.Load() + if lastDecodedNS == 0 { + return + } + now := time.Now() + silentFor := now.Sub(time.Unix(0, lastDecodedNS)) + if silentFor < inputStallSilentThreshold { + return + } + // Only arm the watchdog if the server has been actively sending video + // packets recently. A genuinely static screen (server sends nothing) is + // not a decoder stall and should not trigger a force-break. + if recvTime := g.h264dec.LastReceiveTime(); recvTime.IsZero() || + time.Since(recvTime) >= inputStallSilentThreshold { + return + } + + inputNS := now.UnixNano() + g.inputWatchdogMu.Lock() + g.inputWatchdogNS = inputNS + if g.inputWatchdog == nil { + g.inputWatchdog = time.AfterFunc(localInputRecoveryGrace, func() { + g.fireInputWatchdog(inputNS) + }) + } else { + g.inputWatchdog.Reset(localInputRecoveryGrace) + } + g.inputWatchdogMu.Unlock() + + slog.Debug("H.264: local input armed stall watchdog", + "silentFor", silentFor.Round(time.Millisecond)) +} + +func (g *GfxHandler) fireInputWatchdog(inputNS int64) { + g.inputWatchdogMu.Lock() + if g.inputWatchdogNS != inputNS { + g.inputWatchdogMu.Unlock() + return + } + g.inputWatchdog = nil + g.inputWatchdogMu.Unlock() + + select { + case g.watchdogCh <- struct{}{}: + default: + } +} + +func (g *GfxHandler) stopInputWatchdog() { + g.inputWatchdogMu.Lock() + if g.inputWatchdog != nil { + g.inputWatchdog.Stop() + g.inputWatchdog = nil + } + g.inputWatchdogNS = 0 + g.inputWatchdogMu.Unlock() +} + +// startAuxDecoderBrokenTimer arms a one-shot timer. If it fires before +// stopAuxDecoderBrokenTimer cancels it, it signals watchdogCh so that +// decodeLoop calls maybeRenegotiateCapabilities. +// Idempotent: only the first call while h264dec2 is nil takes effect. +func (g *GfxHandler) startAuxDecoderBrokenTimer() { + g.auxDecoderBrokenTimerMu.Lock() + defer g.auxDecoderBrokenTimerMu.Unlock() + if g.auxDecoderBrokenTimer == nil { + g.auxDecoderBrokenTimer = time.AfterFunc(auxDecoderBrokenTimeout, func() { + // Clear the pointer so startAuxDecoderBrokenTimer can re-arm + // the timer for subsequent retry attempts (e.g. after a + // ForceRefresh in case 2 of maybeRenegotiateCapabilities). + g.auxDecoderBrokenTimerMu.Lock() + g.auxDecoderBrokenTimer = nil + g.auxDecoderBrokenTimerMu.Unlock() + select { + case g.watchdogCh <- struct{}{}: + default: + } + }) + } +} + +func (g *GfxHandler) stopAuxDecoderBrokenTimer() { + g.auxDecoderBrokenTimerMu.Lock() + if g.auxDecoderBrokenTimer != nil { + g.auxDecoderBrokenTimer.Stop() + g.auxDecoderBrokenTimer = nil + } + g.auxDecoderBrokenTimerMu.Unlock() +} + +// maybeRenegotiateCapabilities is called from decodeLoop when watchdogCh fires. +// If h264dec2 has been nil since the timer was armed AND the server has been +// actively sending LC=2 frames, we either reconnect (if LC=2 was previously +// working) or degrade gracefully to LC=0 only (if LC=2 never worked this session). +// Reconnecting mid-session when LC=2 never worked would just produce another +// identical cycle, since the server appears to not include stream2 in LC=0 IDRs. +func (g *GfxHandler) maybeRenegotiateCapabilities() { + if g.h264dec2 != nil { + return // aux decoder recovered while timer was in flight + } + if g.decoderBrokenNotified { + return // reconnect already in flight + } + if g.lc2PermanentlyDegraded { + return // already degraded to LC=0 only; buffered watchdog signals must not re-enter + } + // Only act when the server has recently been sending LC=2 frames — + // a genuinely idle server needs no intervention. + lastLC2NS := g.lastLC2RecvTime.Load() + if lastLC2NS == 0 || time.Since(time.Unix(0, lastLC2NS)) >= auxDecoderBrokenTimeout { + // No recent LC=2 activity (server idle, or HW-decoder GOP-boundary + // stall). Reset the retry counter so that when LC=2 resumes the + // next burst gets fresh ForceRefresh attempts rather than inheriting + // a stale count from a previous burst. + g.auxDecoderNoIDRRetries = 0 + return + } + if !g.stream2EverSeen { + // stream2 has never appeared in any LC=0 packet this session. The server + // (e.g. VirtualBox VRDE) does not include stream2 in LC=0 IDRs and does + // not send standalone LC=2 IDR frames, so the aux decoder can never be + // primed. Reconnecting would reproduce the same failure. Degrade + // gracefully: LC=2 is silently skipped for the remainder of this session. + slog.Warn("H.264: server never sent stream2 in LC=0, LC=2 degraded to LC=0 only (no reconnect)") + g.lc2PermanentlyDegraded = true + return + } + if !g.lc2EverDecoded { + // stream2 has been seen in LC=0 packets (server supports LC=2), but + // the aux decoder has not yet produced a combined frame. The server + // may be sending only P-frame stream2 data in LC=0 packets — no IDR + // means primeAuxDecoder can never initialise h264dec2. + // + // Send a ForceRefresh keyframe request on each attempt so the server + // hopefully includes a stream2 IDR in the next LC=0 IDR packet. + // Windows servers typically begin a new GOP every 30–60 s; with + // auxDecoderMaxIDRRetries=5 (×10 s = 50 s) we cover at least one + // natural boundary before falling back to a reconnect. + // + // The retry counter is reset when the early-return above fires (no + // recent LC=2 activity, e.g. HW-decoder stall), so resumed bursts + // always start from attempt 1. + g.auxDecoderNoIDRRetries++ + if g.auxDecoderNoIDRRetries <= auxDecoderMaxIDRRetries { + slog.Debug("H.264: aux decoder not primed — requesting keyframe to get stream2 IDR", + "attempt", g.auxDecoderNoIDRRetries, "maxRetries", auxDecoderMaxIDRRetries) + g.lastKeyframeRequest = time.Time{} // allow immediate send + g.maybeRequestKeyframe() + return + } + // The server has not delivered a stream2 IDR despite repeated ForceRefresh + // requests and LC=2 has never decoded successfully this session. Reconnecting + // reproduces the same failure because the server consistently omits stream2 + // IDRs in LC=0 IDR packets. Degrade gracefully to LC=0-only instead. + slog.Warn("H.264: aux decoder never primed despite keyframe request — degrading to LC=0 only (no reconnect)", + "retries", g.auxDecoderNoIDRRetries) + g.lc2PermanentlyDegraded = true + return + } + // The timer firing is itself the auxDecoderBrokenTimeout signal. Do not + // gate on lastDecodedFrame: the main decoder (h264dec) may still be active + // and continuously updating lastDecodedFrame even while the aux decoder + // (h264dec2) is broken, which would prevent this function from ever + // triggering a reconnect and permanently lose LC=2 chroma enhancement. + slog.Debug("H.264: aux decoder absent, server sending LC=2 — triggering reconnect") + g.decoderBrokenNotified = true + if g.onDecoderBroken != nil { + go g.onDecoderBroken() + } +} + +func (g *GfxHandler) noteSuccessfulDecode() { + g.lastDecodedFrame.Store(time.Now().UnixNano()) + g.noIDRSoftResetCount = 0 + // A single clean decode resets the *consecutive* drop counter, but does NOT + // disarm swFallbackPrimed. After a HW stall the SW decoder is primed with a + // stale cached IDR; the corruption it causes only manifests as the following + // P-frames diverge from the (missing) reference, so the very first frame can + // decode cleanly and then the picture rots into green/zero-chroma output. + // Clearing swFallbackPrimed here would permanently disable + // trackSWFallbackDroppedFrame's escalation and let that corruption persist. + // swFallbackPrimed is instead cleared only on a genuine fresh IDR resync + // (maybeCacheStream1IDR) or on RESET_GRAPHICS. + g.swFallbackDroppedCount = 0 + g.swFallbackFirstDropTime = time.Time{} + g.stopInputWatchdog() +} + +func (g *GfxHandler) maybeTriggerInputStall() { + g.inputWatchdogMu.Lock() + inputNS := g.inputWatchdogNS + g.inputWatchdogNS = 0 + g.inputWatchdog = nil + g.inputWatchdogMu.Unlock() + + if inputNS == 0 || g.h264dec == nil || g.h264dec.IsBroken() { + return + } + // Don't force-break a decoder that is already in IDR-wait state (e.g. + // after a soft reset). It is in a known recovery path; breaking it again + // would restart the soft-reset cycle unnecessarily. + if g.h264dec.NeedsIDR() { + return + } + // Don't force-break when the server has been idle (no video packets + // arriving). A static screen produces no frames but the decoder is + // healthy; only break when packets are flowing but output is absent. + if recvTime := g.h264dec.LastReceiveTime(); recvTime.IsZero() || + time.Since(recvTime) >= inputStallSilentThreshold { + return + } + lastDecodedNS := g.lastDecodedFrame.Load() + if lastDecodedNS == 0 || lastDecodedNS >= inputNS { + return + } + silentFor := time.Since(time.Unix(0, lastDecodedNS)) + if silentFor < inputStallSilentThreshold { + return + } + g.h264dec.ForceBroken(H264BrokenReasonHWStall) + slog.Debug("H.264: local input produced no new frame, marking decoder broken", + "silentFor", silentFor.Round(time.Millisecond), + "inputAgo", time.Since(time.Unix(0, inputNS)).Round(time.Millisecond)) +} + +// SetSendFunc sets the function used to send RDPGFX responses via DVC. +func (g *GfxHandler) SetSendFunc(fn func([]byte)) { + g.sendFn = fn +} + +// SetDecoderBrokenCallback registers a function that is called once when the +// H.264 decoder becomes permanently unrecoverable (all soft resets exhausted). +// The callback should reconnect the RDP session so a fresh decoder can be +// created from scratch. +func (g *GfxHandler) SetDecoderBrokenCallback(fn func()) { + g.onDecoderBroken = fn +} + +// SetKeyframeRequestFunc registers a function that is called after each +// soft decoder reset to ask the server for a fresh IDR keyframe. This +// speeds up recovery: without it the decoder waits for the server to +// spontaneously send a keyframe. A typical implementation calls +// pdu.SendRefreshRect with the current screen dimensions. +func (g *GfxHandler) SetKeyframeRequestFunc(fn func()) { + g.onKeyframeRequest = fn +} + +// SetH264RawCallback registers a function that receives raw H.264 NAL unit +// data when the built-in decoder is unavailable (h264dec == nil). This +// allows the caller to hand off decoding to an external engine such as the +// browser WebCodecs VideoDecoder API. +// +// destX and destY are the top-left canvas coordinates of the decoded frame. +// isKey is true when the NAL data starts a new GOP (IDR frame). +func (g *GfxHandler) SetH264RawCallback(fn func(destX, destY, w, h int, isKey bool, data []byte, regions []int32)) { + g.onH264Raw = fn +} + +// SetI420Callback registers a callback that receives I420 planar data when an +// H.264 frame is decoded and the underlying decoder supports I420 extraction. +// When set, H264 frames are NOT emitted via the normal OnBitmap path; the +// caller is responsible for rendering the I420 data directly (e.g. via an +// SDL2 IYUV texture). When the I420 fast path is used, the BGRA surface +// backing store is not updated for that frame. +// Set fn to nil to disable and revert to normal OnBitmap delivery. +func (g *GfxHandler) SetI420Callback(fn func(destX, destY, w, h int, y []byte, yStride int, u []byte, uStride int, v []byte, vStride int)) { + g.onI420 = fn +} + +// SetNV12Callback registers a callback that receives native NV12 planar data +// when H.264 decoding produces NV12. When set for AVC420 frames, the normal +// OnBitmap path is bypassed for frames that can be delivered as NV12; callers +// should upload the Y and UV planes directly (for example with SDL2 +// SDL_UpdateNVTexture). Set fn to nil to disable. +func (g *GfxHandler) SetNV12Callback(fn func(destX, destY, w, h int, y []byte, yStride int, uv []byte, uvStride int)) { + g.onNV12 = fn +} + +// SetAVC444Disabled controls whether AVC444/AVC444v2 is advertised to the +// server. When disabled, CAPS_ADVERTISE only includes v8.0 and v8.1, so the +// server will encode frames using AVC420 (4:2:0) only and never send LC=2 +// chroma-upgrade data. This avoids the colour degradation caused by servers +// (e.g. VirtualBox VRDE) that send LC=2 frames without including stream2 in +// LC=0 IDR packets. Must be called before the channel is opened. +func (g *GfxHandler) SetAVC444Disabled(v bool) { + g.avc444Disabled = v +} + +// SetAVCDisabled advertises v10.x caps with RDPGFX_CAPS_FLAG_AVC_DISABLED while +// keeping the RDPGFX channel alive: the server then encodes with ClearCodec and +// RFX Progressive only ("RemoteFX mode"), never H.264. Must be called before +// the channel is opened. +func (g *GfxHandler) SetAVCDisabled(v bool) { + g.avcDisabled = v +} + +// SetPduRecorder installs a callback that receives every wire-to-surface +// bitmap payload for offline replay analysis. Pass nil to disable. +func (g *GfxHandler) SetPduRecorder(fn func(kind byte, codecId uint16, surfW, surfH, x, y, w, h uint32, payload []byte)) { + g.pduRecord = fn +} + +// Replay harness record kinds (see SetPduRecorder). +const ( + ReplayKindWireToSurface1 byte = 1 + ReplayKindWireToSurface2 byte = 2 + ReplayKindCacheToSurface byte = 3 + ReplayKindSurfaceToCache byte = 4 + ReplayKindSolidFill byte = 5 + ReplayKindSurfaceToSurf byte = 6 + ReplayKindEvictCache byte = 7 + ReplayKindResetGraphics byte = 8 +) + +// ReplaySurface ensures a replay surface with the given id/dimensions exists. +func (g *GfxHandler) ReplaySurface(id uint16, w, h uint16) { + g.onCreateSurface([]byte{ + byte(id), byte(id >> 8), + byte(w), byte(w >> 8), + byte(h), byte(h >> 8), + 0x20, + }) +} + +// ReplaySurfacePixels returns the live pixel buffer of a replay surface. +func (g *GfxHandler) ReplaySurfacePixels(id uint16) []byte { + if s, ok := g.surfaces[id]; ok { + return s.data + } + return nil +} + +// ReplayPDU feeds one recorded payload through the normal decode path +// (offline replay harness for garbled-screen debugging). +func (g *GfxHandler) ReplayPDU(kind byte, codecId uint16, surfW, surfH uint16, x, y, w, h uint32, payload []byte) { + switch kind { + case ReplayKindWireToSurface1, ReplayKindWireToSurface2: + if kind == ReplayKindWireToSurface1 { + pdu := make([]byte, 17+len(payload)) + binary.LittleEndian.PutUint16(pdu[0:], 0) + binary.LittleEndian.PutUint16(pdu[2:], codecId) + pdu[4] = 0x20 + binary.LittleEndian.PutUint16(pdu[5:], uint16(x)) + binary.LittleEndian.PutUint16(pdu[7:], uint16(y)) + binary.LittleEndian.PutUint16(pdu[9:], uint16(x+w)) + binary.LittleEndian.PutUint16(pdu[11:], uint16(y+h)) + binary.LittleEndian.PutUint32(pdu[13:], uint32(len(payload))) + copy(pdu[17:], payload) + g.onWireToSurface1Decode(pdu) + return + } + pdu := make([]byte, 13+len(payload)) + binary.LittleEndian.PutUint16(pdu[0:], 0) + binary.LittleEndian.PutUint16(pdu[2:], codecId) + binary.LittleEndian.PutUint32(pdu[4:], 0) + pdu[8] = 0x20 + binary.LittleEndian.PutUint32(pdu[9:], uint32(len(payload))) + copy(pdu[13:], payload) + g.onWireToSurface2Decode(pdu) + case ReplayKindCacheToSurface: + g.onCacheToSurface(payload) + case ReplayKindSurfaceToCache: + g.onSurfaceToCache(payload) + case ReplayKindSolidFill: + g.onSolidFill(payload) + case ReplayKindSurfaceToSurf: + g.onSurfaceToSurface(payload) + case ReplayKindEvictCache: + g.onEvictCacheEntry(payload) + case ReplayKindResetGraphics: + g.onResetGraphics(payload) + } +} + +// OnChannelCreated is called after the DVC CREATE_RSP has been sent. +// It sends RESET_GRAPHICS (MS-RDPEGFX 3.2.1.3:通道初始化后、caps 协商前, +// 客户端先发 RESET_GRAPHICS 声明会话尺寸) 再发 CAPS_ADVERTISE 启动管线。 +// 服务器收到后重建表面并重发完整桌面帧——首次连接无害,断线重附会话时 +// 则强制全量重绘,消除重连后的画面残留。 +// 实测 Win10 19041 依赖该 reset 启动图形管线(省略后服务器不发 caps +// confirm 直接断连);Server 2025 对 reset 后的管线重建会崩(0x112F), +// 由上层对该错误做一次性 bitmap 回退兜底。 +func (g *GfxHandler) OnChannelCreated() { + g.sendResetGraphics() + g.sendCapsAdvertise() +} + +// SetSessionSize 注入会话尺寸,供通道初始化的 RESET_GRAPHICS 使用。 +func (g *GfxHandler) SetSessionSize(width, height uint16) { + g.sessionW = uint32(width) + g.sessionH = uint32(height) +} + +// sendResetGraphics sends RDPGFX_RESET_GRAPHICS_PDU: width(4) + height(4) + +// monitorCount(4)=1,负载共 12 字节。 +// 注意:规范 2.2.2.10 要求整个 PDU 补齐到 340 字节(含 20 字节/显示器的 +// MONITOR_DEF 数组),但实测 Win10 19041 只接受这个 12 字节负载的短格式 +// (340 字节规范格式会令其静默拒绝 GFX 通道回退传统位图;全零 MONITOR_DEF +// 同样被拒)。Server 2025 对任何 reset 格式都会崩(0x112F),由前端 +// ERRINFO_GRAPHICS_* 一次性 bitmap 回退兜底,见 app.js scheduleReconnect。 +func (g *GfxHandler) sendResetGraphics() { + w, h := g.sessionW, g.sessionH + if w == 0 || h == 0 { + return + } + payload := make([]byte, 12) + binary.LittleEndian.PutUint32(payload[0:], w) + binary.LittleEndian.PutUint32(payload[4:], h) + binary.LittleEndian.PutUint32(payload[8:], 1) + g.sendPdu(cmdidResetGraphics, payload) +} + +// sendCapsAdvertise sends RDPGFX_CAPS_ADVERTISE_PDU to the server. +// The client must advertise its capabilities before the server will +// send any graphics data (MS-RDPEGFX 2.2.3.1). +func (g *GfxHandler) sendCapsAdvertise() { + p := pduBufPool.Get().([]byte)[:0] + defer pduBufPool.Put(p[:0]) + + // AVC capsets are advertised when we can deliver decoded frames either + // in-process (h264dec) or by handing the raw NALs off to the embedder + // (onH264Raw, used by the WASM build to forward to WebCodecs). Without + // either, the v8.0+AVCDisabled fallback below forces the server to + // reject RDPGFX and use legacy bitmap PDUs. + if g.h264dec != nil || g.onH264Raw != nil { + if g.avc444Disabled { + // AVC444 关闭(AVC420-only):v8.0/v8.1 之外再广告 v10.0—— + // Windows 服务器在 v10+ caps 下才启用 H.264 编码;只到 v10.0 + //(不含 v10.1+)即不暗示 AVC444v2(LC=2 色度升级流),服务器 + // 最高使用 AVC420。 + p = binary.LittleEndian.AppendUint16(p, 3) // capsSetCount + + // v8.0 — baseline fallback (no AVC) + p = binary.LittleEndian.AppendUint32(p, capVersion8) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagThinClient) + + // v8.1 — AVC420 via explicit flag + p = binary.LittleEndian.AppendUint32(p, capVersion81) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache|capFlagAVC420Enabled) + + // v10.0 — H.264 启用门槛,仅 AVC420 + p = binary.LittleEndian.AppendUint32(p, capVersion10) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache|capFlagAVC420Enabled) + + g.sendPdu(cmdidCapsAdvertise, p) + slog.Debug("RDPGFX: sent CAPS_ADVERTISE (v10.0..v8.0, AVC420 only)") + } else { + // Advertise capsets in ascending order (v8.0 → v10.7), matching + // rdpyqt / FreeRDP layout so servers pick the highest common version. + p = binary.LittleEndian.AppendUint16(p, 11) // capsSetCount + + // v8.0 — baseline fallback (no AVC) + p = binary.LittleEndian.AppendUint32(p, capVersion8) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagThinClient) + + // v8.1 — AVC420 via explicit flag + p = binary.LittleEndian.AppendUint32(p, capVersion81) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache|capFlagAVC420Enabled) + + // v10.0 + p = binary.LittleEndian.AppendUint32(p, capVersion10) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.1 — 16-byte capsData (12 zero bytes after flags) + p = binary.LittleEndian.AppendUint32(p, capVersion101) + p = binary.LittleEndian.AppendUint32(p, 16) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + + // v10.2 + p = binary.LittleEndian.AppendUint32(p, capVersion102) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.3 + p = binary.LittleEndian.AppendUint32(p, capVersion103) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, 0) + + // v10.4 + p = binary.LittleEndian.AppendUint32(p, capVersion104) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.5 + p = binary.LittleEndian.AppendUint32(p, capVersion105) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.6 + p = binary.LittleEndian.AppendUint32(p, capVersion106) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.6.1 + p = binary.LittleEndian.AppendUint32(p, capVersion1061) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.7 + p = binary.LittleEndian.AppendUint32(p, capVersion107) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + g.sendPdu(cmdidCapsAdvertise, p) + slog.Debug("RDPGFX: sent CAPS_ADVERTISE (v10.7..v8.0, AVC enabled)") + } + } else if g.avcDisabled { + // RemoteFX 模式:保留 RDPGFX 通道、显式禁用 AVC——服务端继续用 + // ClearCodec + RFX Progressive 编码(capset 组合与 FreeRDP 无 H264 + // 构建一致)。此前仅广告 v8.0+AVC_DISABLED(该标志自 v10 才在规范 + // 中定义),服务器会拒绝整个 GFX 通道退化成传统 RLE 位图更新, + // 拖动窗口时带宽可达 80Mbps+。 + // v10.3 特意不带 SMALL_CACHE(该版本固定 16MB 缓存槽)。 + noAVC := capFlagSmallCache | capFlagAVCDisabled + p = binary.LittleEndian.AppendUint16(p, 11) // capsSetCount + + // v8.0 / v8.1 — AVC_DISABLED 在 v8 未定义,仅 SMALL_CACHE + p = binary.LittleEndian.AppendUint32(p, capVersion8) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + p = binary.LittleEndian.AppendUint32(p, capVersion81) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagSmallCache) + + // v10.x — SMALL_CACHE | AVC_DISABLED + p = binary.LittleEndian.AppendUint32(p, capVersion10) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + // v10.1 — 16 字节布局(flags + 12 字节保留域) + p = binary.LittleEndian.AppendUint32(p, capVersion101) + p = binary.LittleEndian.AppendUint32(p, 16) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + p = binary.LittleEndian.AppendUint32(p, 0) + + p = binary.LittleEndian.AppendUint32(p, capVersion102) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + p = binary.LittleEndian.AppendUint32(p, capVersion103) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, capFlagAVCDisabled) + + p = binary.LittleEndian.AppendUint32(p, capVersion104) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + p = binary.LittleEndian.AppendUint32(p, capVersion105) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + p = binary.LittleEndian.AppendUint32(p, capVersion106) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + p = binary.LittleEndian.AppendUint32(p, capVersion1061) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + p = binary.LittleEndian.AppendUint32(p, capVersion107) + p = binary.LittleEndian.AppendUint32(p, 4) + p = binary.LittleEndian.AppendUint32(p, noAVC) + + g.sendPdu(cmdidCapsAdvertise, p) + slog.Info("RDPGFX: sent CAPS_ADVERTISE (v10.7..v8.0, AVC disabled, RFX/Clear only)") + } else { + // 传统位图回退模式:v8.0 + AVC_DISABLED 使服务器"拒绝 RDPGFX 通道", + // 退回传统 RLE 位图更新。适用于图形管线损坏(任何 RDPGFX 会话都会 + // 触发 ERRINFO 0x112F 崩溃)或不支持 H264 的服务器。 + // 注意:AVC_DISABLED 只在 v8 caps 上安全(服务器会忽略非法组合)。 + p = binary.LittleEndian.AppendUint16(p, 1) // capsSetCount + p = binary.LittleEndian.AppendUint32(p, capVersion8) + p = binary.LittleEndian.AppendUint32(p, 4) // capsDataLength + p = binary.LittleEndian.AppendUint32(p, capFlagThinClient|capFlagSmallCache|capFlagAVCDisabled) + g.sendPdu(cmdidCapsAdvertise, p) + slog.Debug("RDPGFX: sent CAPS_ADVERTISE (v8.0, AVC disabled → legacy bitmap fallback)") + } +} + +// ZGFX segment descriptors (MS-RDPEGFX 2.2.4) +const ( + zgfxSingle = 0xE0 + zgfxMultipart = 0xE1 + + zgfxCompressedRDP8 = 0x04 +) + +// Process handles a complete RDPGFX payload (may contain multiple PDUs). +// Data arrives wrapped in ZGFX (RDP8 Bulk Compression) segments (MS-RDPEGFX 2.2.4). +// +// Called on the network read goroutine. Decompression happens here; +// the decompressed payload is then queued for asynchronous processing +// (including frame ACKs and decode) on the decode goroutine. +// This keeps the read goroutine free from any socket.Write calls that +// could cause TCP deadlock when both sides try to write simultaneously. +func (g *GfxHandler) Process(data []byte) { + defer func() { + if r := recover(); r != nil { + slog.Error("RDPGFX: panic in Process", "err", r) + } + }() + if len(data) < 1 { + return + } + + var decompressed []byte + var decompPooled bool + + descriptor := data[0] + switch descriptor { + case zgfxSingle: + if len(data) < 2 { + return + } + decompressed, decompPooled = g.decompressSegment(data[1:]) + case zgfxMultipart: + decompressed, decompPooled = g.decompressMultipart(data[1:]) + default: + slog.Warn("RDPGFX: unknown ZGFX descriptor", "descriptor", descriptor) + decompressed = data + } + + if len(decompressed) == 0 { + return + } + + // decompressSegment / decompressMultipart already return owned + // buffers (freshly allocated or copied from input) so we can hand + // the slice directly to the async decode goroutine. + pkt := decodePkt{data: decompressed, pooled: decompPooled} + // 阻塞式入队,绝不丢弃(协议正确性):被丢帧的区域服务器不会重发 + //(FrameAck 已视为送达),表现为永久黑块/花屏。满时靠背压阻塞读 + // 协程 → TCP 流控自然减速服务器;配合 onEndFrame 的 SUSPEND ACK + //(队列深时显式请求服务器暂停发送,排空后真实 queueDepth 恢复), + // 把内存占用约束在 decodeCh 容量级别。 + g.decodeCh <- pkt +} + +// ackDroppedFrames scans decompressed PDU data for EndFrame commands +// and sends ACKs for them. Called on the read goroutine when decodeCh +// is full and the message is being dropped. Without this, dropped +// EndFrames would leave the server's outstanding-frame count stuck, +// eventually causing it to stop sending entirely. +// +// queueDepth is set to suspendFrameAcknowledge (0xFFFFFFFF) so the +// server suspends sending new frames. As the decodeLoop drains the +// existing queue and sends ACKs with lower queueDepth values, the server +// will automatically resume (MS-RDPEGFX 2.2.2.8). +func (g *GfxHandler) ackDroppedFrames(pkt decodePkt) { + data := pkt.data + defer func() { + if pkt.pooled { + releaseBitmapBuf(data) + } + }() + for offset := 0; offset+headerSize <= len(data); { + cmdId := binary.LittleEndian.Uint16(data[offset:]) + pduLength := binary.LittleEndian.Uint32(data[offset+4:]) + if pduLength < uint32(headerSize) || int(pduLength) > len(data)-offset { + break + } + if cmdId == cmdidEndFrame { + pduData := data[offset+headerSize : offset+int(pduLength)] + if len(pduData) >= 4 { + g.sendFrameAck(binary.LittleEndian.Uint32(pduData), suspendFrameAcknowledge) + } + } + offset += int(pduLength) + } +} + +// decompressSegment handles a single ZGFX segment (after the descriptor byte). +// First byte is RDP8_BULK_ENCODED_DATA header: +// +// bits 0-3: compression type (0x04 = RDP8) +// bit 5: PACKET_COMPRESSED (0x20) +// +// Always returns pooled=true: the returned slice was acquired from +// bitmapBufPool (directly or via Decompress) and must be released with +// releaseBitmapBuf once the caller is done with it. +func (g *GfxHandler) decompressSegment(seg []byte) ([]byte, bool) { + if len(seg) < 1 { + return nil, false + } + header := seg[0] + payload := seg[1:] + if header&0x20 != 0 { + // Acquire a pool buffer as the initial backing for Decompress output. + // Decompress may grow beyond it; the returned slice (not buf) must be + // released by the caller. Any over-small buf that gets replaced is + // abandoned to GC — the pool converges to the right size over time. + buf := acquireBitmapBuf(len(payload) * 3) + return g.zgfx.Decompress(payload, buf), true + } + g.zgfx.historyWrite(payload) + // Return a pooled copy: payload aliases the caller's network buffer, which + // will be reused on the next read. Callers hand the slice off to the async + // decode goroutine and must own the memory. + buf := acquireBitmapBuf(len(payload)) + copy(buf, payload) + return buf, true +} + +// decompressMultipart handles ZGFX multipart segments and returns the +// concatenated decompressed data (without processing PDUs). +// Returns a slice acquired from bitmapBufPool; caller must release it. +func (g *GfxHandler) decompressMultipart(data []byte) ([]byte, bool) { + if len(data) < 6 { + return nil, false + } + // Direct slice indexing — avoids bytes.NewReader and per-field io.ReadFull. + segCount := binary.LittleEndian.Uint16(data[0:]) + uncompSize := binary.LittleEndian.Uint32(data[2:]) + offset := 6 + + // Pre-allocate to the advertised uncompressed size to avoid repeated + // buffer growths as each segment is appended. + buf := acquireBitmapBuf(int(uncompSize)) + result := buf[:0] + for range segCount { + if offset+4 > len(data) { + break + } + segSize := int(binary.LittleEndian.Uint32(data[offset:])) + offset += 4 + if offset+segSize > len(data) { + break + } + segData := data[offset : offset+segSize] + offset += segSize + raw, rawPooled := g.decompressSegment(segData) + if raw != nil { + result = append(result, raw...) + if rawPooled { + releaseBitmapBuf(raw) + } + } + } + if len(result) == 0 { + releaseBitmapBuf(buf) + return nil, false + } + // If result grew beyond buf, buf was abandoned; result is the new owner. + return result, true +} + +// decodeLoop runs in a dedicated goroutine, reading decompressed PDU data +// from decodeCh and dispatching all processing — including frame ACKs and +// heavy decode work. Keeping socket.Write calls off the read goroutine +// avoids TCP deadlock (where both sides try to write while neither reads). +// It automatically restarts on panic, unless Close() has been called. +// +// decodeLoop owns the h264dec lifecycle: it is the sole caller of Decode() +// and it frees h264dec on final exit (when doneCh is closed) to avoid a +// use-after-free race with Close() freeing the AVCodecContext from a +// different goroutine while Decode() holds it. +func (g *GfxHandler) decodeLoop() { + defer func() { + if r := recover(); r != nil { + slog.Error("RDPGFX: panic in decodeLoop, restarting", "err", r, "stack", string(debug.Stack())) + select { + case <-g.doneCh: + // Shutting down — free h264dec resources on this final exit. + if g.h264dec != nil { + g.h264dec.Close() + g.h264dec = nil + } + if g.h264dec2 != nil { + g.h264dec2.Close() + g.h264dec2 = nil + } + default: + go g.decodeLoop() + } + return + } + // Normal exit triggered by doneCh being closed. + if g.h264dec != nil { + g.h264dec.Close() + g.h264dec = nil + } + if g.h264dec2 != nil { + g.h264dec2.Close() + g.h264dec2 = nil + } + }() + slog.Debug("RDPGFX: decodeLoop started") + for { + select { + case <-g.doneCh: + return + case pkt := <-g.decodeCh: + g.decodePDUs(pkt.data) + if pkt.pooled { + releaseBitmapBuf(pkt.data) + } + case <-g.watchdogCh: + // A background timer in ffmpegDecoder fired because the stall-probe + // or IDR-wait timeout expired while the server was sending no frames + // (static screen → near-0 fps), or because local input produced no + // new frame while the decoder had already been silent too long. + slog.Debug("H.264: watchdog triggered, checking decoder state") + g.maybeTriggerInputStall() + g.maybeRenegotiateCapabilities() + g.maybeNotifyDecoderBroken() + } + } +} + +// decodePDUs processes all PDUs in decompressed data. +// Every PDU is decoded: silently dropping CaVideo/progressive updates under +// backpressure permanently diverges the canvas (the server considers acked +// frames delivered and never resends them — the corruption manifests as +// persistent noise after a fast window drag). Pacing is the server's job +// via the queueDepth we report in FRAME_ACKNOWLEDGE. +func (g *GfxHandler) decodePDUs(data []byte) { + start := time.Now() + defer func() { + // 单条消息解码耗时的 EMA(α=0.3):诊断解码吞吐是否跟不上到达速率 + if cost := time.Since(start).Microseconds(); cost > 0 { + for { + cur := g.decUs.Load() + next := cur * 7 / 10 + if next == 0 { + next = cost + } else { + next += cost * 3 / 10 + } + if g.decUs.CompareAndSwap(cur, next) { + return + } + } + } + }() + for offset := 0; offset+headerSize <= len(data); { + cmdId := binary.LittleEndian.Uint16(data[offset:]) + pduLength := binary.LittleEndian.Uint32(data[offset+4:]) + if pduLength < uint32(headerSize) || int(pduLength) > len(data)-offset { + break + } + pduData := data[offset+headerSize : offset+int(pduLength)] + g.dispatchDecode(cmdId, pduData) + offset += int(pduLength) + } +} + +// dispatchDecode routes a single PDU. +func (g *GfxHandler) dispatchDecode(cmdId uint16, data []byte) { + if int(cmdId) < len(g.cmdCounts) { + g.cmdCounts[cmdId].Add(1) + } + switch cmdId { + case cmdidCapsConfirm: + g.onCapsConfirm(data) + case cmdidResetGraphics: + g.onResetGraphics(data) + case cmdidCreateSurface: + g.onCreateSurface(data) + case cmdidDeleteSurface: + g.onDeleteSurface(data) + case cmdidMapSurfaceToOutput: + g.onMapSurfaceToOutput(data) + case cmdidStartFrame: + // 记录帧解码起点,供 QoE 周期上报(timeDiffSE)使用 + g.frameDecodeStart.Store(time.Now().UnixNano()) + case cmdidSurfaceToSurface: + g.onSurfaceToSurface(data) + case cmdidSurfaceToCache: + g.onSurfaceToCache(data) + case cmdidEndFrame: + g.onEndFrame(data) // always ACK, even under backpressure + case cmdidWireToSurface1: + g.onWireToSurface1Decode(data) + case cmdidWireToSurface2: + g.onWireToSurface2Decode(data) + case cmdidSolidFill: + g.onSolidFill(data) + case cmdidCacheToSurface: + g.onCacheToSurface(data) + case cmdidEvictCacheEntry: + g.onEvictCacheEntry(data) + case cmdidCacheImportOffer: + g.onCacheImportOffer() + case cmdidMapSurfaceToWindow, cmdidMapSurfaceToScaledWindow: + // ignored — we don't support per-window mapping + case cmdidCacheImportReply: + g.onCacheImportReply(data) + case cmdidDeleteEncodingContext, cmdidQoeFrameAcknowledge: + // no client state to maintain for these + case cmdidMapSurfaceToScaledOutput: + g.onMapSurfaceToScaledOutput(data) + default: + slog.Debug("RDPGFX: unhandled cmd", "cmdId", cmdId) + } +} + +// writeLoop runs in a dedicated goroutine. It reads serialized ACK +// PDUs from ackCh and sends each one via sendFn. Every ACK must reach +// the server — the server tracks outstanding frames individually and +// stops sending if ACKs are missing. Automatically restarts on panic, +// unless Close() has been called. +func (g *GfxHandler) writeLoop() { + defer func() { + if r := recover(); r != nil { + slog.Error("RDPGFX: panic in writeLoop, restarting", "err", r) + select { + case <-g.doneCh: + // Shut down; do not restart. + default: + go g.writeLoop() + } + } + }() + for { + select { + case <-g.doneCh: + return + case pdu := <-g.ackCh: + if g.sendFn != nil { + g.sendFn(pdu) + } + ackPDUPool.Put(pdu) + } + } +} + +// sendPdu sends a PDU synchronously. Used for rare control messages +// (CapsAdvertise, CacheImportReply) that must not be dropped. +// pduBufPool reuses scratch byte slices for assembling outbound PDU frames, +// avoiding per-call heap allocations on the sendPdu hot path. +var pduBufPool = sync.Pool{ + New: func() any { return make([]byte, 0, headerSize+256) }, +} + +// ackPDUPool reuses the fixed-size 20-byte slices used for FRAME_ACKNOWLEDGE +// PDUs (~60/s during video), eliminating per-frame heap allocations. +var ackPDUPool = sync.Pool{ + New: func() any { return make([]byte, 20) }, +} + +func (g *GfxHandler) sendPdu(cmdId uint16, payload []byte) { + if g.sendFn == nil { + return + } + buf := pduBufPool.Get().([]byte) + buf = buf[:0] + buf = binary.LittleEndian.AppendUint16(buf, cmdId) + buf = binary.LittleEndian.AppendUint16(buf, 0) // flags + buf = binary.LittleEndian.AppendUint32(buf, uint32(headerSize+len(payload))) + buf = append(buf, payload...) + g.sendFn(buf) + pduBufPool.Put(buf[:0]) +} + +// --- Command Handlers --- + +// DebugSurfacePixel 诊断:返回第一个 mapped surface 上 (x,y) 的 BGRA 原始值 +func (g *GfxHandler) DebugSurfacePixel(x, y int) (uint8, uint8, uint8, uint8, bool) { + for _, s := range g.surfaces { + if !s.mapped { + continue + } + if x < 0 || y < 0 || x >= int(s.width) || y >= int(s.height) { + return 0, 0, 0, 0, false + } + i := (y*int(s.width) + x) * 4 + return s.data[i], s.data[i+1], s.data[i+2], s.data[i+3], true + } + return 0, 0, 0, 0, false +} + +func (g *GfxHandler) onCapsConfirm(data []byte) { + if len(data) < 12 { + slog.Debug("RDPGFX: CAPS_CONFIRM received (short)") + return + } + version := binary.LittleEndian.Uint32(data[0:]) + dataLen := binary.LittleEndian.Uint32(data[4:]) + flags := uint32(0) + if dataLen >= 4 { + flags = binary.LittleEndian.Uint32(data[8:]) + } + slog.Info("RDPGFX: CAPS_CONFIRM", "version", fmt.Sprintf("0x%08X", version), "flags", fmt.Sprintf("0x%08X", flags)) + // caps 交换完成后上报持久缓存条目(MS-RDPEGFX 2.2.2.16)。此时连接的 + // 其他 GFX 状态尚未开始流动,是最安全的上报时点。 + g.sendCacheImportOffer() +} + +func (g *GfxHandler) onResetGraphics(data []byte) { + if len(data) < 12 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindResetGraphics, 0, 0, 0, 0, 0, 0, 0, data) + } + w := binary.LittleEndian.Uint32(data[0:]) + h := binary.LittleEndian.Uint32(data[4:]) + // RESET_GRAPHICS 意味着服务端图形子系统重启(常见于 0x112F 内部错误), + // 后续帧流是否恢复依赖此处的状态清理,值得始终留痕。 + slog.Info("RDPGFX: RESET_GRAPHICS", "w", w, "h", h) + g.surfaces = make(map[uint16]*surface) + g.clearCtx = newClearCodecCtx() + g.framesDecoded.Store(0) + g.softResetCount = 0 + g.noIDRSoftResetCount = 0 + g.decoderBrokenNotified = false + g.lc2EverDecoded = false + g.stream2EverSeen = false + g.auxDecoderNoIDRRetries = 0 + g.lc2PermanentlyDegraded = false + g.lastKeyframeRequest = time.Time{} + g.lastStream1IDR = g.lastStream1IDR[:0] + g.lastStream1IDRTime = time.Time{} + g.lastStream1IDRFrame = 0 + g.swFallbackPrimed = false + g.swFallbackDroppedCount = 0 + g.swFallbackFirstDropTime = time.Time{} + g.lastDecodedFrame.Store(0) + g.stopInputWatchdog() + if g.h264dec != nil { + g.h264dec.Close() + g.h264dec = newH264DecoderWithWatchdog(g.watchdogCh) + } + if g.h264dec2 != nil { + g.h264dec2.Close() + // Keep h264dec2 nil; primeAuxDecoder will recreate it on the next stream2 IDR + // so the fresh decoder is always primed before receiving LC=2 P-frames. + g.h264dec2 = nil + } + g.avc444YPlane = avc444YPlane{} + g.avc444IDRYPlane = avc444YPlane{} + g.progressive.Reset() +} + +func (g *GfxHandler) onCreateSurface(data []byte) { + if len(data) < 7 { + return + } + id := binary.LittleEndian.Uint16(data[0:]) + w := binary.LittleEndian.Uint16(data[2:]) + h := binary.LittleEndian.Uint16(data[4:]) + f := data[6] + slog.Debug("RDPGFX: CREATE_SURFACE", "id", id, "w", w, "h", h) + g.surfaces[id] = &surface{ + width: w, height: h, format: f, + data: make([]byte, int(w)*int(h)*4), + shadowStale: true, + } +} + +func (g *GfxHandler) onDeleteSurface(data []byte) { + if len(data) < 2 { + return + } + id := binary.LittleEndian.Uint16(data) + delete(g.surfaces, id) +} + +func (g *GfxHandler) onMapSurfaceToOutput(data []byte) { + if len(data) < 12 { + return + } + id := binary.LittleEndian.Uint16(data[0:]) + // data[2:4] = reserved + ox := binary.LittleEndian.Uint32(data[4:]) + oy := binary.LittleEndian.Uint32(data[8:]) + slog.Debug("RDPGFX: MAP_SURFACE", "id", id, "ox", ox, "oy", oy) + if s, ok := g.surfaces[id]; ok { + s.outputX = ox + s.outputY = oy + s.mapped = true + } +} + +func (g *GfxHandler) onMapSurfaceToScaledOutput(data []byte) { + if len(data) < 20 { + return + } + id := binary.LittleEndian.Uint16(data[0:]) + // data[2:4] = reserved + ox := binary.LittleEndian.Uint32(data[4:]) + oy := binary.LittleEndian.Uint32(data[8:]) + // data[12:16] = targetWidth, data[16:20] = targetHeight (unused) + slog.Debug("RDPGFX: MAP_SURFACE_SCALED", "id", id, "ox", ox, "oy", oy) + if s, ok := g.surfaces[id]; ok { + s.outputX = ox + s.outputY = oy + s.mapped = true + } +} + +// sendFrameAck builds and queues a FRAME_ACKNOWLEDGE PDU. +// Safe to call from any goroutine (uses atomic framesDecoded). +// The PDU is serialized directly into a 20-byte slice to avoid +// the two bytes.Buffer allocations the previous implementation required. +// +// queueDepth is reported to the server so it can adjust encoding quality +// and frame rate based on the client's decode backlog. Pass +// suspendFrameAcknowledge (0xFFFFFFFF) to ask the server to suspend new +// frames until a subsequent ACK with a lower value is received. +func (g *GfxHandler) sendFrameAck(frameId uint32, queueDepth uint32) { + decoded := g.framesDecoded.Add(1) + // 8-byte RDPGFX header + 12-byte FRAME_ACKNOWLEDGE payload = 20 bytes. + pdu := ackPDUPool.Get().([]byte) + binary.LittleEndian.PutUint16(pdu[0:], cmdidFrameAcknowledge) + // pdu[2:4] = flags (0) — zero value + binary.LittleEndian.PutUint16(pdu[2:], 0) + binary.LittleEndian.PutUint32(pdu[4:], 20) // total PDU length + binary.LittleEndian.PutUint32(pdu[8:], queueDepth) + binary.LittleEndian.PutUint32(pdu[12:], frameId) + binary.LittleEndian.PutUint32(pdu[16:], decoded) + select { + case g.ackCh <- pdu: + default: + ackPDUPool.Put(pdu) + slog.Warn("RDPGFX: ackCh full, ACK dropped") + } +} + +func (g *GfxHandler) onEndFrame(data []byte) { + if len(data) < 4 { + return + } + // 帧间隔 EMA(α=0.3):诊断服务器帧率与传输节奏 + now := time.Now().UnixNano() + if prev := g.lastFrameAt.Swap(now); prev != 0 { + intv := (now - prev) / 1000 // µs + for { + cur := g.frameIntvUs.Load() + next := cur * 7 / 10 + if next == 0 { + next = intv + } else { + next += intv * 3 / 10 + } + if g.frameIntvUs.CompareAndSwap(cur, next) { + break + } + } + } + realDepth := uint32(len(g.decodeCh)) + if hint := g.queueDepthHint.Load(); hint > realDepth { + realDepth = hint + } + frameId := binary.LittleEndian.Uint32(data) + // 背压:队列深时显式 SUSPEND(queueDepth=0xFFFFFFFF 请求服务器暂停 + // 发送),排空到低水位后以真实 queueDepth 恢复(MS-RDPEGFX 2.2.2.8)。 + // 这是协议规定的限速通道——配合读协程的阻塞式入队,保证任何已发送 + // 的帧都会被完整解码上屏,不存在"ACK 了却没画"的黑块来源。 + switch depth := uint32(len(g.decodeCh)); { + case depth > 48: + g.suspended.Store(true) + case depth <= 16: + g.suspended.Store(false) + } + if g.suspended.Load() { + g.sendFrameAck(frameId, suspendFrameAcknowledge) + return + } + g.sendFrameAck(frameId, realDepth) + // QoE 周期上报(每 30 帧):mstsc/FreeRDP 用它向服务器反馈解码时延, + // 服务器据此做码率/画质自适应。cmdId=0x16,布局对齐 FreeRDP + // rdpgfx_send_qoe_frame_acknowledge_pdu(header8 + frameId4 + + // timestamp4 + timeDiffSE2 + timeDiffEDR2 = 20 字节,与 ack 同尺寸, + // 复用 ackCh 写循环)。 + if frameId%30 == 0 { + now := time.Now().UnixNano() + startNs := g.frameDecodeStart.Load() + if startNs == 0 { + startNs = now + } + diffSE := (now - startNs) / int64(time.Millisecond) + if diffSE < 0 { + diffSE = 0 + } + if diffSE > 65000 { + diffSE = 65000 + } + edr := g.decUs.Load() / 1000 + if edr > 65000 { + edr = 65000 + } + pdu := ackPDUPool.Get().([]byte) + binary.LittleEndian.PutUint16(pdu[0:], cmdidQoeFrameAcknowledge) + binary.LittleEndian.PutUint16(pdu[2:], 0) + binary.LittleEndian.PutUint32(pdu[4:], 20) + binary.LittleEndian.PutUint32(pdu[8:], frameId) + binary.LittleEndian.PutUint32(pdu[12:], uint32(startNs/int64(time.Millisecond))) + binary.LittleEndian.PutUint16(pdu[16:], uint16(diffSE)) + binary.LittleEndian.PutUint16(pdu[18:], uint16(edr)) + select { + case g.ackCh <- pdu: + default: + ackPDUPool.Put(pdu) + } + } +} + +// DiagStats 返回实时诊断指标:累计解码帧数、当前解码队列深度、 +// 帧间隔(毫秒)与单消息解码耗时 EMA(微秒)。解码耗时必须保留 +// 微秒精度——WebCodecs 路径单消息只有几百微秒,取整到毫秒恒为 0。 +// 统计循环定期输出,用于判断解码吞吐是否跟不上到达速率、服务器帧率是否异常。 +func (g *GfxHandler) DiagStats() (frames, qdepth, fintvMs, decUs int64) { + return int64(g.framesDecoded.Load()), + int64(len(g.decodeCh)), + g.frameIntvUs.Load() / 1000, + g.decUs.Load() +} + +// SetQueueDepthHint sets a minimum queueDepth to report in FRAME_ACKNOWLEDGE +// PDUs (MS-RDPEGFX 2.2.2.8). The server uses this value to pace its frame +// rate and encoding quality: a larger value signals that the client's decode +// queue is full, causing the server to slow down or reduce quality. +// +// A hint of 0 (the default) means "report the real queue length". +// Values in the range 10–100 are typical for moderate throttling. +// Use suspendFrameAcknowledge (0xFFFFFFFF) to pause the stream entirely +// (the stream resumes automatically when the hint is cleared). + +// CodecStats returns cumulative surface-bitmap bytes per codec id +// (WTS1 + WTS2 combined). Useful for bandwidth diagnostics: e.g. a session +// dominated by codec 3 (RemoteFX) vs 0x0B/0x0E (AVC420/444) vs 8 (ClearCodec). +// Keys 100+cmdId carry the per-command PDU counters (100+0..100+63). +func (g *GfxHandler) CodecStats() map[uint16]int64 { + out := make(map[uint16]int64, len(g.codecBytes)) + for i := range g.codecBytes { + if v := g.codecBytes[i].Load(); v != 0 { + out[uint16(i)] = v + } + } + for i := range g.cmdCounts { + if v := g.cmdCounts[i].Load(); v != 0 { + out[uint16(100+i)] = v + } + } + return out +} + +func (g *GfxHandler) SetQueueDepthHint(depth uint32) { + g.queueDepthHint.Store(depth) +} + +// onWireToSurface1Decode handles RDPGFX_WIRE_TO_SURFACE_PDU_1 (MS-RDPEGFX 2.2.2.1). +func (g *GfxHandler) onWireToSurface1Decode(data []byte) { + if len(data) < 17 { + return + } + // Parse fixed header fields via direct binary indexing (avoids bytes.NewReader + // and per-field io.ReadFull overhead on the hot H.264 path). + surfId := binary.LittleEndian.Uint16(data[0:]) + codecId := binary.LittleEndian.Uint16(data[2:]) + pixFmt := data[4] + left := binary.LittleEndian.Uint16(data[5:]) + top := binary.LittleEndian.Uint16(data[7:]) + right := binary.LittleEndian.Uint16(data[9:]) + bottom := binary.LittleEndian.Uint16(data[11:]) + bmpLen := binary.LittleEndian.Uint32(data[13:]) + if int(bmpLen) > len(data)-17 { + return + } + // 服务端在拖动时会夹杂 pixFmt 非法(0x62 等)的垃圾微型更新, + // FreeRDP 对非法 pixFmt 直接拒绝 PDU;照单全收会在屏幕上画出 + // 彩色杂线(拖动花屏的组成部分)。 + if pixFmt != pixelFormatXRGB8888 && pixFmt != pixelFormatARGB8888 { + return + } + bmpData := data[17 : 17+int(bmpLen)] + + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: WTS1", "surfId", surfId, "codecId", codecId, + "w", right-left, "h", bottom-top, "bmpLen", bmpLen) + } + + w := int(right - left) + h := int(bottom - top) + if w <= 0 || h <= 0 { + return + } + if int(codecId) < len(g.codecBytes) { + g.codecBytes[codecId].Add(int64(bmpLen)) + } + + s, ok := g.surfaces[surfId] + if !ok { + return + } + if g.pduRecord != nil { + g.pduRecord(1, codecId, uint32(s.width), uint32(s.height), + uint32(left), uint32(top), uint32(w), uint32(h), bmpData) + } + + // CaVideo (0x0003) carries RFX tile-encoded data; decode onto the + // persistent surface buffer like the progressive codec in WTS2. + if codecId == codecCaVideo { + rects := g.rfx.Decode(bmpData, int(left), int(top), s.data, int(s.width), int(s.height)) + g.emitCaVideoRects(s, rects) + return + } + + var decoded []byte + var avcRegions []avcRect + owned := false // true ⇒ decoded buffer is from bitmapBufPool and must be released + switch codecId { + case codecUncompressed: + decoded = decodeUncompressed(bmpData, w, h, pixFmt) + owned = true + case codecPlanar: + decoded = decodePlanar(bmpData, w, h) + owned = true + case codecAVC420: + destX := int(s.outputX) + int(left) + destY := int(s.outputY) + int(top) + if g.onNV12 != nil { + var ownedAVC bool + var nv12 *H264FrameNV12 + decoded, nv12, avcRegions, ownedAVC = g.decodeAVC420WithNV12(bmpData, destX, destY, w, h) + owned = ownedAVC + if nv12 != nil { + if decoded != nil { + blitToSurface(s, int(left), int(top), w, h, decoded) + if owned { + releaseBitmapBuf(decoded) + } + } else { + // Display advanced without a shadow update; mark stale so a + // later full-surface (WTS2) frame fully repairs the shadow. + s.shadowStale = true + } + g.onNV12(destX, destY, w, h, nv12.Y, nv12.YStride, nv12.UV, nv12.UVStride) + return + } + // NV12 unavailable; fall through to BGRA emit. + } else if g.onI420 != nil { + var ownedAVC bool + var i420 *H264FrameI420 + decoded, i420, avcRegions, ownedAVC = g.decodeAVC420WithI420(bmpData, destX, destY, w, h) + owned = ownedAVC + if i420 != nil { + if decoded != nil { + blitToSurface(s, int(left), int(top), w, h, decoded) + if owned { + releaseBitmapBuf(decoded) + } + } else { + s.shadowStale = true + } + g.onI420(destX, destY, w, h, i420.Y, i420.YStride, i420.U, i420.UStride, i420.V, i420.VStride) + return + } + // I420 unavailable (nil frame or unsupported format); fall through to BGRA emit. + } else { + var ownedAVC bool + decoded, avcRegions, ownedAVC = g.decodeAVC420(bmpData, destX, destY, w, h) + owned = ownedAVC + } + case codecAVC444, codecAVC444v2: + destX := int(s.outputX) + int(left) + destY := int(s.outputY) + int(top) + if g.onNV12 != nil { + var ownedAVC bool + var nv12 *H264FrameNV12 + decoded, nv12, avcRegions, ownedAVC = g.decodeAVC444WithNV12(bmpData, destX, destY, w, h) + owned = ownedAVC + if nv12 != nil { + if decoded != nil { + blitToSurface(s, int(left), int(top), w, h, decoded) + if owned { + releaseBitmapBuf(decoded) + } + } else { + // Display advanced without a shadow update; mark stale so a + // later full-surface frame fully repairs the shadow. + s.shadowStale = true + } + g.onNV12(destX, destY, w, h, nv12.Y, nv12.YStride, nv12.UV, nv12.UVStride) + return + } + // nv12 == nil: LC=2 chroma frame or decoder stall. + // decoded may contain combined BGRA; fall through to emit it. + } else if g.onI420 != nil { + var ownedAVC bool + var i420 *H264FrameI420 + decoded, i420, avcRegions, ownedAVC = g.decodeAVC444WithI420(bmpData, destX, destY, w, h) + owned = ownedAVC + if i420 != nil { + if decoded != nil { + blitToSurface(s, int(left), int(top), w, h, decoded) + if owned { + releaseBitmapBuf(decoded) + } + } else { + s.shadowStale = true + } + g.onI420(destX, destY, w, h, i420.Y, i420.YStride, i420.U, i420.UStride, i420.V, i420.VStride) + return + } + // i420 == nil: LC=2 or decoder unavailable; decoded may contain BGRA. + } else { + var ownedAVC bool + decoded, avcRegions, ownedAVC = g.decodeAVC444(bmpData, destX, destY, w, h) + owned = ownedAVC + } + case codecClear: + // ClearCodec 承载 UI/文字/图标等内容(现代服务器主力编码器), + // 解码器维护跨帧 vBar 缓存,输出 BGRA 位图 + decoded = g.clearCtx.decode(bmpData, w, h) + owned = true + case codecProgressive: + // Progressive 瓦片按 surface 网格绝对定位,直接解到持久 surface 缓冲 + //(与 WTS2 分支相同)。服务器可能经任一 WTS 消息投递 progressive。 + rects := g.progressive.Decode(bmpData, s.data, int(s.width), int(s.height)) + stride := int(s.width) * 4 + for _, rc := range rects { + needed := rc.w * rc.h * 4 + region := regionPool.Get().([]byte) + if cap(region) < needed { + region = make([]byte, needed) + } else { + region = region[:needed] + } + rowBytes := rc.w * 4 + for row := 0; row < rc.h; row++ { + srcOff := (rc.y+row)*stride + rc.x*4 + dstOff := row * rowBytes + if srcOff+rowBytes <= len(s.data) { + copy(region[dstOff:dstOff+rowBytes], s.data[srcOff:srcOff+rowBytes]) + } + } + g.emitBitmap(s, rc.x, rc.y, rc.w, rc.h, region) + regionPool.Put(region) + } + return + default: + slog.Warn("RDPGFX: unsupported codec in WTS1", "codecId", codecId, "surfId", surfId, "w", w, "h", h, "bmpLen", bmpLen) + return + } + if decoded == nil { + return + } + + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, int(left), int(top), w, h, decoded, avcRegions) + if owned { + releaseBitmapBuf(decoded) + } + return + } + + blitToSurface(s, int(left), int(top), w, h, decoded) + if owned { + g.emitBitmapPooled(s, int(left), int(top), w, h, decoded) + } else { + g.emitBitmap(s, int(left), int(top), w, h, decoded) + } +} + +// onWireToSurface2Decode handles RDPGFX_WIRE_TO_SURFACE_PDU_2 (MS-RDPEGFX 2.2.2.2). +func (g *GfxHandler) onWireToSurface2Decode(data []byte) { + if len(data) < 13 { + return + } + // Parse fixed header fields via direct binary indexing (avoids bytes.NewReader + // and per-field io.ReadFull overhead on the hot H.264 path). + surfId := binary.LittleEndian.Uint16(data[0:]) + codecId := binary.LittleEndian.Uint16(data[2:]) + codecCtxId := binary.LittleEndian.Uint32(data[4:]) + pixFmt := data[8] + bmpLen := binary.LittleEndian.Uint32(data[9:]) + if int(bmpLen) > len(data)-13 { + return + } + if pixFmt != pixelFormatXRGB8888 && pixFmt != pixelFormatARGB8888 { + return + } + bmpData := data[13 : 13+int(bmpLen)] + + s, ok := g.surfaces[surfId] + if !ok { + return + } + + w := int(s.width) + h := int(s.height) + + if g.pduRecord != nil { + g.pduRecord(2, codecId, uint32(s.width), uint32(s.height), 0, 0, + uint32(w), uint32(h), bmpData) + } + if slog.Default().Enabled(nil, slog.LevelDebug) { + slog.Debug("RDPGFX: WTS2", "surfId", surfId, "codecId", codecId, + "w", w, "h", h, "bmpLen", bmpLen) + } + if int(codecId) < len(g.codecBytes) { + g.codecBytes[codecId].Add(int64(bmpLen)) + } + + var decoded []byte + switch codecId { + case codecUncompressed: + decoded = decodeUncompressed(bmpData, w, h, pixFmt) + blitToSurface(s, 0, 0, w, h, decoded) + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + case codecPlanar: + decoded = decodePlanar(bmpData, w, h) + blitToSurface(s, 0, 0, w, h, decoded) + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + case codecClear: + clearDecoded := g.clearCtx.decode(bmpData, w, h) + blitToSurface(s, 0, 0, w, h, clearDecoded) + g.emitBitmapPooled(s, 0, 0, w, h, clearDecoded) + case codecCaVideo: + rects := g.rfx.Decode(bmpData, 0, 0, s.data, w, h) + g.emitCaVideoRects(s, rects) + case codecAVC420: + destX := int(s.outputX) + destY := int(s.outputY) + if g.onNV12 != nil { + decoded, nv12, avcRegions, ownedAVC := g.decodeAVC420WithNV12(bmpData, destX, destY, w, h) + if nv12 != nil { + if decoded != nil { + // GPU drives the display via onNV12; the CPU shadow only needs + // the dirty regions when it is already in sync. A full blit is + // forced when the shadow is stale (a prior GPU-only frame + // advanced the display without updating it). + if !s.shadowStale && len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAVCRegionsToSurface(s, 0, 0, w, h, decoded, avcRegions) + } else { + blitToSurface(s, 0, 0, w, h, decoded) + s.shadowStale = false + } + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + // Display advanced (onNV12) without a BGRA shadow update; mark + // the shadow stale so the next decoded frame fully repairs it. + s.shadowStale = true + } + g.onNV12(destX, destY, w, h, nv12.Y, nv12.YStride, nv12.UV, nv12.UVStride) + } else if decoded != nil { + // NV12 unavailable; fall back to BGRA emit. + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, 0, 0, w, h, decoded, avcRegions) + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + blitToSurface(s, 0, 0, w, h, decoded) + if ownedAVC { + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + } else { + g.emitBitmap(s, 0, 0, w, h, decoded) + } + } + } + } else if g.onI420 != nil { + decoded, i420, avcRegions, ownedAVC := g.decodeAVC420WithI420(bmpData, destX, destY, w, h) + if i420 != nil { + if decoded != nil { + if !s.shadowStale && len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAVCRegionsToSurface(s, 0, 0, w, h, decoded, avcRegions) + } else { + blitToSurface(s, 0, 0, w, h, decoded) + s.shadowStale = false + } + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + s.shadowStale = true + } + g.onI420(destX, destY, w, h, i420.Y, i420.YStride, i420.U, i420.UStride, i420.V, i420.VStride) + } else if decoded != nil { + // I420 unavailable; fall back to BGRA emit. + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, 0, 0, w, h, decoded, avcRegions) + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + blitToSurface(s, 0, 0, w, h, decoded) + if ownedAVC { + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + } else { + g.emitBitmap(s, 0, 0, w, h, decoded) + } + } + } + } else { + decoded, avcRegions, ownedAVC := g.decodeAVC420(bmpData, destX, destY, w, h) + if decoded != nil { + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, 0, 0, w, h, decoded, avcRegions) + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + blitToSurface(s, 0, 0, w, h, decoded) + if ownedAVC { + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + } else { + g.emitBitmap(s, 0, 0, w, h, decoded) + } + } + } + } + case codecAVC444, codecAVC444v2: + destX := int(s.outputX) + destY := int(s.outputY) + if g.onI420 != nil { + decoded, i420, avcRegions, ownedAVC := g.decodeAVC444WithI420(bmpData, destX, destY, w, h) + if i420 != nil { + if decoded != nil { + if !s.shadowStale && len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAVCRegionsToSurface(s, 0, 0, w, h, decoded, avcRegions) + } else { + blitToSurface(s, 0, 0, w, h, decoded) + s.shadowStale = false + } + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + s.shadowStale = true + } + g.onI420(destX, destY, w, h, i420.Y, i420.YStride, i420.U, i420.UStride, i420.V, i420.VStride) + } else if decoded != nil { + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, 0, 0, w, h, decoded, avcRegions) + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + blitToSurface(s, 0, 0, w, h, decoded) + if ownedAVC { + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + } else { + g.emitBitmap(s, 0, 0, w, h, decoded) + } + } + } + } else { + decoded, avcRegions, ownedAVC := g.decodeAVC444(bmpData, destX, destY, w, h) + if decoded != nil { + if len(avcRegions) > 0 && shouldUseAVCRegions(avcRegions, w, h) { + g.blitAndEmitAVCRegions(s, 0, 0, w, h, decoded, avcRegions) + if ownedAVC { + releaseBitmapBuf(decoded) + } + } else { + blitToSurface(s, 0, 0, w, h, decoded) + if ownedAVC { + g.emitBitmapPooled(s, 0, 0, w, h, decoded) + } else { + g.emitBitmap(s, 0, 0, w, h, decoded) + } + } + } + } + case codecProgressive: + // Decode tiles directly onto the persistent surface buffer. + rects := g.progressive.Decode(bmpData, s.data, w, h) + for _, rc := range rects { + needed := rc.w * rc.h * 4 + region := regionPool.Get().([]byte) + if cap(region) < needed { + region = make([]byte, needed) + } else { + region = region[:needed] + } + stride := w * 4 + rowBytes := rc.w * 4 + for row := 0; row < rc.h; row++ { + srcOff := (rc.y+row)*stride + rc.x*4 + dstOff := row * rowBytes + if srcOff+rowBytes <= len(s.data) { + copy(region[dstOff:dstOff+rowBytes], s.data[srcOff:srcOff+rowBytes]) + } + } + g.emitBitmap(s, rc.x, rc.y, rc.w, rc.h, region) + regionPool.Put(region) + } + default: + slog.Debug("RDPGFX: WTS2 unsupported codec", "codecId", codecId, "ctxId", codecCtxId) + return + } +} + +func (g *GfxHandler) onSolidFill(data []byte) { + if len(data) < 8 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindSolidFill, 0, 0, 0, 0, 0, 0, 0, data) + } + surfId := binary.LittleEndian.Uint16(data[0:]) + cb := data[2] + cg := data[3] + cr := data[4] + // data[5] = XA (ignored) + fillCount := binary.LittleEndian.Uint16(data[6:]) + + s, ok := g.surfaces[surfId] + if !ok { + return + } + + stride := int(s.width) * 4 + // Pre-compose a single BGRA pixel as a uint32 for one-shot writes. + pixelU32 := uint32(cb) | uint32(cg)<<8 | uint32(cr)<<16 | uint32(0xFF)<<24 + + offset := 8 + for range fillCount { + if offset+8 > len(data) { + break + } + left := binary.LittleEndian.Uint16(data[offset:]) + top := binary.LittleEndian.Uint16(data[offset+2:]) + right := binary.LittleEndian.Uint16(data[offset+4:]) + bottom := binary.LittleEndian.Uint16(data[offset+6:]) + offset += 8 + w := int(right - left) + h := int(bottom - top) + if w <= 0 || h <= 0 { + continue + } + + // Clamp to surface bounds + yEnd := min(int(bottom), int(s.height)) + xEnd := min(int(right), int(s.width)) + + // Fill the first row with PutUint32 (single 32-bit store per pixel), + // then replicate it to subsequent rows with copy(). + rowStart := int(top)*stride + int(left)*4 + rowBytes := (xEnd - int(left)) * 4 + if rowStart+rowBytes <= len(s.data) { + row := s.data[rowStart : rowStart+rowBytes] + for x := 0; x+4 <= rowBytes; x += 4 { + binary.LittleEndian.PutUint32(row[x:], pixelU32) + } + for y := int(top) + 1; y < yEnd; y++ { + dst := y*stride + int(left)*4 + if dst+rowBytes <= len(s.data) { + copy(s.data[dst:dst+rowBytes], row) + } + } + } + + if s.mapped && g.onBitmap != nil { + // Build fill data: fill first row, then replicate (doubling). + fillData := acquireBitmapBuf(w * h * 4) + rowW := w * 4 + for x := 0; x+4 <= rowW; x += 4 { + binary.LittleEndian.PutUint32(fillData[x:], pixelU32) + } + // Doubling copy: O(log h) memmoves instead of h linear copies. + filled := rowW + total := rowW * h + for filled*2 <= total { + copy(fillData[filled:filled*2], fillData[:filled]) + filled *= 2 + } + if filled < total { + copy(fillData[filled:total], fillData[:total-filled]) + } + destL := int(s.outputX) + int(left) + destT := int(s.outputY) + int(top) + g.singleUpdate[0] = BitmapUpdate{ + DestLeft: destL, DestTop: destT, + DestRight: destL + w - 1, DestBottom: destT + h - 1, + Width: w, Height: h, Bpp: 4, Data: fillData, + } + g.emitAndReleaseUpdates(g.singleUpdate[:]) + } + } +} + +// onSurfaceToSurface handles RDPGFX_SURFACE_TO_SURFACE_PDU (MS-RDPEGFX +// 2.2.2.5). Layout: srcId(2) dstId(2) rectSrc(8) destPtsCount(2) destPts[] +// (4 bytes each). rectSrc on the source surface is copied to each +// destination point — this is how the server scrolls/blits content during +// window drags, and same-surface copies may overlap, so the row order must +// adapt to the copy direction. +func (g *GfxHandler) onSurfaceToSurface(data []byte) { + // Minimum payload: srcId(2)+dstId(2)+rectSrc(8)+destPtsCount(2) = 14 bytes + if len(data) < 14 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindSurfaceToSurf, 0, 0, 0, 0, 0, 0, 0, data) + } + srcId := binary.LittleEndian.Uint16(data[0:]) + dstId := binary.LittleEndian.Uint16(data[2:]) + left := int(binary.LittleEndian.Uint16(data[4:])) + top := int(binary.LittleEndian.Uint16(data[6:])) + right := int(binary.LittleEndian.Uint16(data[8:])) + bottom := int(binary.LittleEndian.Uint16(data[10:])) + destCount := int(binary.LittleEndian.Uint16(data[12:])) + + src, srcOk := g.surfaces[srcId] + dst, dstOk := g.surfaces[dstId] + if !srcOk || !dstOk { + return + } + + w := right - left + h := bottom - top + if w <= 0 || h <= 0 { + return + } + + srcStride := int(src.width) * 4 + dstStride := int(dst.width) * 4 + rowBytes := w * 4 + offset := 14 + for i := 0; i < destCount; i++ { + if offset+4 > len(data) { + break + } + dstX := int(binary.LittleEndian.Uint16(data[offset:])) + dstY := int(binary.LittleEndian.Uint16(data[offset+2:])) + offset += 4 + if left+rowBytes/4 > int(src.width) || top+h > int(src.height) || + dstX+w > int(dst.width) || dstY+h > int(dst.height) { + continue + } + + // 同面滚动拷贝:目标在源下方时自下而上逐行复制,避免未读源行 + // 被提前覆盖(行内水平重叠由 copy 的 memmove 语义保证)。 + first, last, step := 0, h, 1 + if src == dst && dstY > top { + first, last, step = h-1, -1, -1 + } + for row := first; row != last; row += step { + s := (top+row)*srcStride + left*4 + d := (dstY+row)*dstStride + dstX*4 + copy(dst.data[d:d+rowBytes], src.data[s:s+rowBytes]) + } + // emitBitmap 的 Data 必须是 w×h 的矩形位图:把拷贝结果抽出来再发。 + // 直接传整面 dst.data 会让 JS 端按 w×h 读表面左上角内容贴到 + // 目标位置——拖动时窗口内部被桌面角落内容覆盖,即花屏条纹。 + rectBuf := acquireBitmapBuf(w * h * 4) + for row := 0; row < h; row++ { + srcOff := (dstY+row)*dstStride + dstX*4 + dstOff := row * rowBytes + copy(rectBuf[dstOff:dstOff+rowBytes], dst.data[srcOff:srcOff+rowBytes]) + } + g.emitBitmapPooled(dst, dstX, dstY, w, h, rectBuf) + } +} + +// onSurfaceToCache copies a surface region into a cache slot +// (MS-RDPEGFX 2.2.2.6). Without this the cache stays empty and every +// subsequent CACHE_TO_SURFACE blits nothing — a major source of missing +// UI elements on Win11 servers, which lean on the cache heavily. +func (g *GfxHandler) onSurfaceToCache(data []byte) { + if len(data) < 20 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindSurfaceToCache, 0, 0, 0, 0, 0, 0, 0, data) + } + surfId := binary.LittleEndian.Uint16(data[0:]) + // 持久缓存键:跨重连的身份标识,CacheImportOffer 上报它而非槽位 + cacheKey := binary.LittleEndian.Uint64(data[2:]) + cacheSlot := binary.LittleEndian.Uint16(data[10:]) + left := int(binary.LittleEndian.Uint16(data[12:])) + top := int(binary.LittleEndian.Uint16(data[14:])) + right := int(binary.LittleEndian.Uint16(data[16:])) + bottom := int(binary.LittleEndian.Uint16(data[18:])) + + s, ok := g.surfaces[surfId] + if !ok { + return + } + w := right - left + h := bottom - top + if w <= 0 || h <= 0 || left < 0 || top < 0 || + right > int(s.width) || bottom > int(s.height) { + return + } + + rowBytes := w * 4 + entry := cacheEntry{key: cacheKey, width: w, height: h, data: make([]byte, w*h*4)} + for row := 0; row < h; row++ { + srcOff := (top+row)*int(s.width)*4 + left*4 + dstOff := row * rowBytes + copy(entry.data[dstOff:dstOff+rowBytes], s.data[srcOff:srcOff+rowBytes]) + } + g.cacheEntries[cacheSlot] = entry + // 交浏览器侧持久化:重连时以 CacheImportOffer 上报该键(MS-RDPEGFX + // 持久位图缓存)。key=0 视为服务器未提供有效键,跳过。 + if g.cacheStore != nil && cacheKey != 0 { + g.cacheStore.Persist(cacheKey, w, h, 32, entry.data) + } +} + +func (g *GfxHandler) onCacheToSurface(data []byte) { + if len(data) < 6 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindCacheToSurface, 0, 0, 0, 0, 0, 0, 0, data) + } + cacheSlot := binary.LittleEndian.Uint16(data[0:]) + surfId := binary.LittleEndian.Uint16(data[2:]) + destCount := binary.LittleEndian.Uint16(data[4:]) + + ce, hasCE := g.cacheEntries[cacheSlot] + s, hasSurf := g.surfaces[surfId] + + offset := 6 + for range destCount { + if offset+4 > len(data) { + break + } + dx := binary.LittleEndian.Uint16(data[offset:]) + dy := binary.LittleEndian.Uint16(data[offset+2:]) + offset += 4 + if hasCE && hasSurf { + blitToSurface(s, int(dx), int(dy), ce.width, ce.height, ce.data) + g.emitBitmap(s, int(dx), int(dy), ce.width, ce.height, ce.data) + } + } +} + +func (g *GfxHandler) onEvictCacheEntry(data []byte) { + if len(data) < 2 { + return + } + if g.pduRecord != nil { + g.pduRecord(ReplayKindEvictCache, 0, 0, 0, 0, 0, 0, 0, data) + } + slot := binary.LittleEndian.Uint16(data) + delete(g.cacheEntries, slot) +} + +func (g *GfxHandler) onCacheImportOffer() { + var p [2]byte // importedEntriesCount = 0 (little-endian zero) + g.sendPdu(cmdidCacheImportReply, p[:]) +} + +// maxCacheImportEntries 限制单条 CacheImportOffer 的条目数。规范上限 5461 +// (cacheEntriesCount < 0x1556),静态桌面的瓦片数百量级即够,同时把 +// 浏览器侧存储与导出开销控制在几 MB 内。 +const maxCacheImportEntries = 512 + +// sendCacheImportOffer 在 caps 确认后向服务器上报客户端持久缓存中仍持有的 +// 条目(cacheKey + bitmapLength)。服务器按前缀导入并在 CacheImportReply +// 中回分配的槽位;导入的条目随后可被 CacheToSurface 直接回贴,静态内容 +// 在重连后无需重传(MS-RDPEGFX 2.2.2.16/2.2.2.17)。 +func (g *GfxHandler) sendCacheImportOffer() { + if g.importOfferSent || g.cacheStore == nil { + return + } + g.importOfferSent = true + entries := g.cacheStore.Export() + if len(entries) == 0 { + return + } + if len(entries) > maxCacheImportEntries { + entries = entries[:maxCacheImportEntries] + } + payload := make([]byte, 2, 2+12*len(entries)) + binary.LittleEndian.PutUint16(payload, uint16(len(entries))) + for _, e := range entries { + payload = binary.LittleEndian.AppendUint64(payload, e.Key) + payload = binary.LittleEndian.AppendUint32(payload, uint32(len(e.Data))) + } + g.offeredCache = entries + g.sendPdu(cmdidCacheImportOffer, payload) + slog.Info("RDPGFX: cache import offer", "entries", len(entries)) +} + +// onCacheImportReply 处理 CacheImportOffer 的应答:importedEntriesCount 为 +// 前缀语义——上报条目的前 N 条被导入,cacheSlots[i] 是第 i 条的新槽位。 +// 像素数据客户端本就持有,直接重建槽位映射即可。 +func (g *GfxHandler) onCacheImportReply(data []byte) { + if len(data) < 2 { + return + } + n := int(binary.LittleEndian.Uint16(data)) + if n > len(g.offeredCache) { + n = len(g.offeredCache) + } + if len(data) < 2+2*n { // 槽位数组被截断:只取完整部分 + n = (len(data) - 2) / 2 + } + imported := 0 + for i := 0; i < n; i++ { + e := g.offeredCache[i] + // 任何尺寸/长度不一致的条目绝不入缓存——错误回贴就是花屏 + if e.Width <= 0 || e.Height <= 0 || len(e.Data) != e.Width*e.Height*4 { + continue + } + slot := binary.LittleEndian.Uint16(data[2+i*2:]) + g.cacheEntries[slot] = cacheEntry{key: e.Key, width: e.Width, height: e.Height, data: e.Data} + imported++ + } + g.offeredCache = nil + slog.Info("RDPGFX: cache import reply", "imported", imported, "of", n) +} + +// SetPersistentCacheStore 安装持久缓存桥;必须在连接建立前调用。 +func (g *GfxHandler) SetPersistentCacheStore(s GfxCacheStore) { + g.cacheStore = s +} + +// --- Helpers --- + +// emitCaVideoRects copies decoded RemoteFX tile regions from the surface +// pixel buffer into individual BitmapUpdate slices and emits them. +// Used by both onWireToSurface1Decode and onWireToSurface2Decode. +func (g *GfxHandler) emitCaVideoRects(s *surface, rects []rfxRect) { + if !s.mapped || g.onBitmap == nil || len(rects) == 0 { + return + } + g.updatesBuf = g.updatesBuf[:0] + stride := int(s.width) * 4 + for _, rc := range rects { + needed := rc.w * rc.h * 4 + region := acquireBitmapBuf(needed) + rowBytes := rc.w * 4 + for row := 0; row < rc.h; row++ { + srcOff := (rc.y+row)*stride + rc.x*4 + dstOff := row * rowBytes + if srcOff+rowBytes <= len(s.data) { + copy(region[dstOff:dstOff+rowBytes], s.data[srcOff:srcOff+rowBytes]) + } + } + destL := int(s.outputX) + rc.x + destT := int(s.outputY) + rc.y + g.updatesBuf = append(g.updatesBuf, BitmapUpdate{ + DestLeft: destL, DestTop: destT, + DestRight: destL + rc.w - 1, DestBottom: destT + rc.h - 1, + Width: rc.w, Height: rc.h, Bpp: 4, Data: region, + }) + } + g.emitAndReleaseUpdates(g.updatesBuf) +} + +func blitToSurface(s *surface, x, y, w, h int, src []byte) { + stride := int(s.width) * 4 + // Full-width fast path: when x==0 and w==surface.width the entire region + // is contiguous in both src and s.data — replace h row-copies with one. + if x == 0 && w == int(s.width) && y >= 0 && y+h <= int(s.height) { + dstOff := y * stride + n := h * stride + if n <= len(src) && dstOff+n <= len(s.data) { + copy(s.data[dstOff:dstOff+n], src[:n]) + return + } + } + for row := range h { + dy := y + row + if dy < 0 || dy >= int(s.height) { + continue + } + srcOff := row * w * 4 + dstOff := dy*stride + x*4 + n := w * 4 + if dstOff >= 0 && dstOff+n <= len(s.data) && srcOff+n <= len(src) { + copy(s.data[dstOff:dstOff+n], src[srcOff:srcOff+n]) + } + } +} + +// emitBitmapPooled is like emitBitmap but releases `decoded` back to +// bitmapBufPool after the synchronous onBitmap callback returns. Use this +// for codec output buffers that the GfxHandler owns end-to-end (currently +// uncompressed and planar). +func (g *GfxHandler) emitBitmapPooled(s *surface, x, y, w, h int, decoded []byte) { + if !s.mapped || g.onBitmap == nil { + releaseBitmapBuf(decoded) + return + } + destL := int(s.outputX) + x + destT := int(s.outputY) + y + g.singleUpdate[0] = BitmapUpdate{ + DestLeft: destL, DestTop: destT, + DestRight: destL + w - 1, DestBottom: destT + h - 1, + Width: w, Height: h, Bpp: 4, Data: decoded, + } + g.emitAndReleaseUpdates(g.singleUpdate[:]) +} + +func (g *GfxHandler) emitBitmap(s *surface, x, y, w, h int, decoded []byte) { + if !s.mapped || g.onBitmap == nil { + return + } + destL := int(s.outputX) + x + destT := int(s.outputY) + y + g.singleUpdate[0] = BitmapUpdate{ + DestLeft: destL, DestTop: destT, + DestRight: destL + w - 1, DestBottom: destT + h - 1, + Width: w, Height: h, Bpp: 4, Data: decoded, + } + g.onBitmap(g.singleUpdate[:]) + g.singleUpdate[0].Data = nil // release reference; decoded is not pooled +} + +// --- Codec: Uncompressed --- + +func decodeUncompressed(data []byte, w, h int, pixFmt uint8) []byte { + out := acquireBitmapBuf(w * h * 4) + n := w * h * 4 + if len(data) >= n { + copy(out, data[:n]) + } else { + copy(out[:len(data)], data) + // Zero the unfilled tail in case the slice was reused from the pool. + clear(out[len(data):n]) + } + return out +} + +// --- Codec: Planar (RDP 6.0 Bitmap Codec, MS-RDPEGDI 2.2.2.5) --- + +func decodePlanar(data []byte, w, h int) []byte { + if len(data) < 1 { + return acquireBitmapBuf(w * h * 4) + } + header := data[0] + // 头部标志位(FreeRDP include/freerdp/codec/planar.h): + // bit3 CS 色度子采样、bit4 RLE、bit5 NA(无 alpha 平面)、bits0-2 CLL。 + rle := (header >> 4) & 1 + noAlpha := (header >> 5) & 1 + planeSize := w * h + offset := 1 + + var alphaPlane, redPlane, greenPlane, bluePlane []byte + defer func() { + releasePlaneBuf(alphaPlane) + releasePlaneBuf(redPlane) + releasePlaneBuf(greenPlane) + releasePlaneBuf(bluePlane) + }() + if rle == 0 { + if noAlpha == 0 { + alphaPlane, offset = readRawPlane(data, offset, planeSize) + } + redPlane, offset = readRawPlane(data, offset, planeSize) + greenPlane, offset = readRawPlane(data, offset, planeSize) + bluePlane, offset = readRawPlane(data, offset, planeSize) + } else { + if noAlpha == 0 { + alphaPlane, offset = decodePlanarRLEPlane(data, offset, w, h) + } + redPlane, offset = decodePlanarRLEPlane(data, offset, w, h) + greenPlane, offset = decodePlanarRLEPlane(data, offset, w, h) + bluePlane, offset = decodePlanarRLEPlane(data, offset, w, h) + } + _ = offset + + out := acquireBitmapBuf(planeSize * 4) + // Hoist the per-pixel nil/length checks: clamp each plane to + // `planeSize` (zero-fill missing planes) so the inner loop has no + // branches and the bounds checks are eliminated. + rp := planeOrZero(redPlane, planeSize) + gp := planeOrZero(greenPlane, planeSize) + bp := planeOrZero(bluePlane, planeSize) + ap := alphaPlane + hasAlpha := ap != nil && len(ap) >= planeSize + if hasAlpha { + ap = ap[:planeSize] + for i := range planeSize { + j := i * 4 + out[j] = bp[i] + out[j+1] = gp[i] + out[j+2] = rp[i] + out[j+3] = ap[i] + } + } else { + for i := range planeSize { + j := i * 4 + out[j] = bp[i] + out[j+1] = gp[i] + out[j+2] = rp[i] + out[j+3] = 0xFF + } + } + return out +} + +// planeOrZero returns a slice of exactly `size` bytes, either the input +// plane (truncated if longer) or a zero-filled buffer when the plane is +// nil or short. Used to drop per-pixel nil/bounds checks in decodePlanar. +func planeOrZero(plane []byte, size int) []byte { + if len(plane) >= size { + return plane[:size] + } + out := make([]byte, size) + copy(out, plane) + return out +} + +func readRawPlane(data []byte, offset, size int) ([]byte, int) { + plane := acquirePlaneBuf(size) + end := min(offset+size, len(data)) + if offset < end { + copy(plane, data[offset:end]) + if end-offset < size { + clear(plane[end-offset:]) + } + } else { + clear(plane) + } + return plane, offset + size +} + +// decodePlanarRLEPlane 解码一个 RLE 压缩的像素平面 +// (FreeRDP planar_decompress_plane_rle)。控制字节:低 4 位 runLen、 +// 高 4 位 rawLen;runLen==1 → run=raw+16、runLen==2 → run=raw+32 +// (此时 rawLen 清零)。首行携带绝对像素值,后续行携带相对上一行 +// 同列像素的有符号 delta(最低位为符号位),行程重复最后一个像素 +// 的 delta。 +func decodePlanarRLEPlane(data []byte, offset, w, h int) ([]byte, int) { + out := acquirePlaneBuf(w * h) + clamp := func(v int) byte { + if v < 0 { + return 0 + } + if v > 255 { + return 255 + } + return byte(v) + } + for y := 0; y < h; y++ { + rowStart := y * w + prevRow := rowStart - w + for x := 0; x < w; { + if offset >= len(data) { + return out, offset + } + ctrl := data[offset] + offset++ + runLen := int(ctrl & 0x0F) + rawLen := int((ctrl >> 4) & 0x0F) + switch runLen { + case 1: + runLen = rawLen + 16 + rawLen = 0 + case 2: + runLen = rawLen + 32 + rawLen = 0 + } + + if y == 0 { + // 首行:绝对像素值 + var last byte + for i := 0; i < rawLen; i++ { + if offset >= len(data) || x+i >= w { + return out, offset + } + last = data[offset] + out[rowStart+x+i] = last + offset++ + } + x += rawLen + for i := 0; i < runLen && x < w; i++ { + out[rowStart+x] = last + x++ + } + continue + } + + // 后续行:相对上一行的有符号 delta + lastDelta := 0 + for i := 0; i < rawLen; i++ { + if offset >= len(data) || x+i >= w { + return out, offset + } + dv := data[offset] + offset++ + if dv&1 != 0 { + lastDelta = -int(dv>>1) - 1 + } else { + lastDelta = int(dv >> 1) + } + out[rowStart+x+i] = clamp(int(out[prevRow+x+i]) + lastDelta) + } + x += rawLen + // 行程延续最后一个 delta(相对各自位置上一行的像素) + for i := 0; i < runLen && x < w; i++ { + out[rowStart+x] = clamp(int(out[prevRow+x]) + lastDelta) + x++ + } + } + } + return out, offset +} diff --git a/plugin/rdpgfx/rdpgfx_cache_test.go b/plugin/rdpgfx/rdpgfx_cache_test.go new file mode 100644 index 0000000..33f1111 --- /dev/null +++ b/plugin/rdpgfx/rdpgfx_cache_test.go @@ -0,0 +1,234 @@ +package rdpgfx + +import ( + "bytes" + "encoding/binary" + "testing" +) + +// fakeStore 记录 Persist 调用并回放固定条目,用于验证持久缓存桥。 +type fakeStore struct { + persisted []GfxCacheEntry + export []GfxCacheEntry +} + +func (s *fakeStore) Persist(key uint64, w, h int, bpp uint16, data []byte) { + cp := make([]byte, len(data)) + copy(cp, data) + s.persisted = append(s.persisted, GfxCacheEntry{Key: key, Width: w, Height: h, Bpp: bpp, Data: cp}) +} + +func (s *fakeStore) Export() []GfxCacheEntry { return s.export } + +func (s *fakeStore) Get(key uint64) (GfxCacheEntry, bool) { + for _, e := range s.export { + if e.Key == key { + return e, true + } + } + return GfxCacheEntry{}, false +} + +func (s *fakeStore) Keys() []uint64 { + out := make([]uint64, 0, len(s.export)) + for _, e := range s.export { + out = append(out, e.Key) + } + return out +} + +// newCacheTestHandler 返回带捕获 sendFn 的最小处理器(不经 NewGfxHandler, +// 避免拉起解码/写循环 goroutine)。 +func newCacheTestHandler() (*GfxHandler, *[][]byte) { + sent := &[][]byte{} + g := &GfxHandler{ + surfaces: make(map[uint16]*surface), + cacheEntries: make(map[uint16]cacheEntry), + sendFn: func(b []byte) { + cp := make([]byte, len(b)) + copy(cp, b) + *sent = append(*sent, cp) + }, + } + return g, sent +} + +func TestSendCacheImportOffer(t *testing.T) { + g, sent := newCacheTestHandler() + st := &fakeStore{export: []GfxCacheEntry{ + {Key: 0x1122334455667788, Width: 8, Height: 2, Data: bytes.Repeat([]byte{0xAB}, 8*2*4)}, + {Key: 0x0102030405060708 >> 0, Width: 4, Height: 4, Data: bytes.Repeat([]byte{0xCD}, 4*4*4)}, + }} + g.cacheStore = st + + // 空上报守卫:未调用 Export 前直接发送应只发一次,重复调用被闩住 + g.sendCacheImportOffer() + if len(*sent) != 1 { + t.Fatalf("期望 1 条 PDU,实得 %d", len(*sent)) + } + g.sendCacheImportOffer() + if len(*sent) != 1 { + t.Fatalf("importOfferSent 闩失效:实得 %d 条", len(*sent)) + } + + pdu := (*sent)[0] + if got := binary.LittleEndian.Uint16(pdu[0:]); got != cmdidCacheImportOffer { + t.Fatalf("cmdId=0x%X,期望 0x%X", got, cmdidCacheImportOffer) + } + wantLen := 8 + 2 + 12*2 + if got := binary.LittleEndian.Uint32(pdu[4:]); int(got) != wantLen { + t.Fatalf("pduLength=%d,期望 %d", got, wantLen) + } + if got := binary.LittleEndian.Uint16(pdu[8:]); got != 2 { + t.Fatalf("cacheEntriesCount=%d,期望 2", got) + } + // 条目 1:key u64 + bitmapLength u32 + if got := binary.LittleEndian.Uint64(pdu[10:]); got != 0x1122334455667788 { + t.Fatalf("entry0 key=0x%X", got) + } + if got := binary.LittleEndian.Uint32(pdu[18:]); got != 8*2*4 { + t.Fatalf("entry0 bitmapLength=%d,期望 %d", got, 8*2*4) + } + if got := binary.LittleEndian.Uint64(pdu[22:]); got != 0x0102030405060708 { + t.Fatalf("entry1 key=0x%X", got) + } + if len(g.offeredCache) != 2 { + t.Fatalf("offeredCache 应保留 2 条待映射,实得 %d", len(g.offeredCache)) + } +} + +func TestSendCacheImportOfferEmpty(t *testing.T) { + g, sent := newCacheTestHandler() + g.cacheStore = &fakeStore{export: nil} + g.sendCacheImportOffer() + if len(*sent) != 0 { + t.Fatalf("空存储不应发送 PDU") + } + g.cacheStore = nil + g.importOfferSent = false + g.sendCacheImportOffer() + if len(*sent) != 0 { + t.Fatalf("无 store 不应发送 PDU") + } +} + +func TestOnCacheImportReply(t *testing.T) { + g, _ := newCacheTestHandler() + e0 := GfxCacheEntry{Key: 0xA, Width: 4, Height: 2, Data: bytes.Repeat([]byte{1}, 4*2*4)} + e1 := GfxCacheEntry{Key: 0xB, Width: 2, Height: 2, Data: bytes.Repeat([]byte{2}, 2*2*4)} + g.offeredCache = []GfxCacheEntry{e0, e1} + + // 前缀语义:前 2 条导入,槽位 7 与 9 + data := make([]byte, 2, 2+4) + binary.LittleEndian.PutUint16(data, 2) + data = binary.LittleEndian.AppendUint16(data, 7) + data = binary.LittleEndian.AppendUint16(data, 9) + g.onCacheImportReply(data) + + ce, ok := g.cacheEntries[7] + if !ok || ce.key != 0xA || ce.width != 4 || ce.height != 2 || !bytes.Equal(ce.data, e0.Data) { + t.Fatalf("槽位 7 条目不符: %+v", ce) + } + ce, ok = g.cacheEntries[9] + if !ok || ce.key != 0xB || !bytes.Equal(ce.data, e1.Data) { + t.Fatalf("槽位 9 条目不符: %+v", ce) + } + if g.offeredCache != nil { + t.Fatalf("Reply 后 offeredCache 应清空") + } +} + +func TestOnCacheImportReplyClamp(t *testing.T) { + g, _ := newCacheTestHandler() + bad := GfxCacheEntry{Key: 0xC, Width: 4, Height: 2, Data: []byte{1, 2, 3}} // 长度不符 + ok1 := GfxCacheEntry{Key: 0xD, Width: 2, Height: 2, Data: bytes.Repeat([]byte{3}, 2*2*4)} + g.offeredCache = []GfxCacheEntry{bad, ok1} + + // n=5 超过上报数(截到 2);条目 0 长度不符必须被拒(防花屏) + data := make([]byte, 2, 2+10) + binary.LittleEndian.PutUint16(data, 5) + for _, s := range []uint16{3, 4} { + data = binary.LittleEndian.AppendUint16(data, s) + } + g.onCacheImportReply(data) + + if _, hit := g.cacheEntries[3]; hit { + t.Fatalf("长度不符的条目不应入缓存") + } + if ce, hit := g.cacheEntries[4]; !hit || ce.key != 0xD { + t.Fatalf("槽位 4 应为有效条目") + } + + // 槽位数组截断:只有 1 个完整槽位 + g2, _ := newCacheTestHandler() + g2.offeredCache = []GfxCacheEntry{ok1, ok1} + short := []byte{2, 0, 6, 0} // n=2 但只有 1 个槽位 + g2.onCacheImportReply(short) + if _, hit := g2.cacheEntries[6]; !hit { + t.Fatalf("截断时应导入完整部分") + } + if len(g2.cacheEntries) != 1 { + t.Fatalf("截断时不应导入缺失槽位,实得 %d 条", len(g2.cacheEntries)) + } +} + +func TestSurfaceToCachePersists(t *testing.T) { + g, _ := newCacheTestHandler() + st := &fakeStore{} + g.cacheStore = st + + // 8×4 表面,每行像素值 = 行号(BGRA 同值) + sw, sh := 8, 4 + sdata := make([]byte, sw*sh*4) + for row := range sh { + for col := 0; col < sw; col++ { + o := (row*sw + col) * 4 + sdata[o], sdata[o+1], sdata[o+2], sdata[o+3] = byte(row), byte(row), byte(row), 0xFF + } + } + g.surfaces[1] = &surface{width: uint16(sw), height: uint16(sh), data: sdata} + + key := uint64(0x1122334455667788) + p := make([]byte, 0, 20) + p = binary.LittleEndian.AppendUint16(p, 1) // surfId + p = binary.LittleEndian.AppendUint64(p, key) // cacheKey + p = binary.LittleEndian.AppendUint16(p, 3) // slot + p = binary.LittleEndian.AppendUint16(p, 2) // left + p = binary.LittleEndian.AppendUint16(p, 1) // top + p = binary.LittleEndian.AppendUint16(p, 6) // right + p = binary.LittleEndian.AppendUint16(p, 3) // bottom + g.onSurfaceToCache(p) + + ce := g.cacheEntries[3] + if ce.width != 4 || ce.height != 2 { + t.Fatalf("缓存条目尺寸 %dx%d,期望 4x2", ce.width, ce.height) + } + if ce.key != key { + t.Fatalf("缓存条目 key=0x%X,期望 0x%X", ce.key, key) + } + // 第一行来自表面第 1 行(值为 1),第二行来自第 2 行(值为 2) + if ce.data[0] != 1 || ce.data[(4*1)*4] != 2 { + t.Fatalf("缓存像素内容不符: [0]=%d [row1]=%d", ce.data[0], ce.data[(4*1)*4]) + } + if len(st.persisted) != 1 { + t.Fatalf("Persist 应被调用 1 次,实得 %d", len(st.persisted)) + } + pv := st.persisted[0] + if pv.Key != key || pv.Width != 4 || pv.Height != 2 || !bytes.Equal(pv.Data, ce.data) { + t.Fatalf("持久化条目不符: %+v", pv) + } + + // key=0 不持久化(视为无效键) + p0 := append([]byte(nil), p...) + binary.LittleEndian.PutUint64(p0[2:], 0) + g.onSurfaceToCache(p0) + if len(st.persisted) != 1 { + t.Fatalf("key=0 不应触发 Persist,实得 %d", len(st.persisted)) + } +} + +func TestMaxCacheImportEntries(t *testing.T) { + if maxCacheImportEntries >= 0x1556 { + t.Fatalf("maxCacheImportEntries=%d 必须小于规范上限 0x1556", maxCacheImportEntries) + } +} diff --git a/plugin/rdpgfx/rfx.go b/plugin/rdpgfx/rfx.go new file mode 100644 index 0000000..eca39d6 --- /dev/null +++ b/plugin/rdpgfx/rfx.go @@ -0,0 +1,321 @@ +package rdpgfx + +// Non-progressive RemoteFX (RFX) codec decoder (MS-RDPRFX). +// Used for RDPGFX_CODECID_CAVIDEO (0x0003) in WIRE_TO_SURFACE_PDU_1. +// +// Block type codes (same numeric values as progressive, different semantics): +// 0xCCC0 WBT_SYNC +// 0xCCC1 WBT_CODEC_VERSIONS +// 0xCCC2 WBT_CHANNELS +// 0xCCC3 WBT_CONTEXT (+ 2-byte codecId/channelId) +// 0xCCC4 WBT_FRAME_BEGIN (+ 2-byte codecId/channelId) +// 0xCCC5 WBT_FRAME_END (+ 2-byte codecId/channelId) +// 0xCCC6 WBT_REGION (+ 2-byte codecId/channelId) +// 0xCCC7 WBT_EXTENSION (+ 2-byte codecId/channelId, contains TILESET) +// +// Tile sub-blocks inside TILESET use CBT_TILE (0xCAC3) with standard 6-byte header. + +import ( + "encoding/binary" + "log/slog" + "runtime" + "sync" +) + +const ( + wbtSync = 0xCCC0 + wbtCodecVersions = 0xCCC1 + wbtChannels = 0xCCC2 + wbtContext = 0xCCC3 + wbtFrameBegin = 0xCCC4 + wbtFrameEnd = 0xCCC5 + wbtRegion = 0xCCC6 + wbtExtension = 0xCCC7 + + cbtRegion = 0xCAC1 + cbtTileset = 0xCAC2 + cbtTile = 0xCAC3 +) + +type rfxTileWork struct { + content []byte +} + +type rfxDecoder struct { + rectsBuf []rfxRect + tilesBuf []rfxTileWork + quantsBuf []rfxQuant +} + +func newRfxDecoder() *rfxDecoder { + return &rfxDecoder{} +} + +// Decode processes non-progressive RFX data, rendering tiles onto the +// provided surface buffer at the given (left, top) offset. +// Returns the bounding rectangles of decoded regions in surface coordinates. +func (d *rfxDecoder) Decode(data []byte, left, top int, surfData []byte, width, height int) []rfxRect { + var rects []rfxRect + var quants []rfxQuant + + offset := 0 + for offset+6 <= len(data) { + blockType := binary.LittleEndian.Uint16(data[offset:]) + blockLen := int(binary.LittleEndian.Uint32(data[offset+2:])) + + if blockLen < 6 || offset+blockLen > len(data) { + break + } + + // Determine content start: blocks 0xCCC3-0xCCC7 have 2 extra bytes + // (codecId + channelId) per TS_RFX_CODEC_CHANNELT. + headerLen := 6 + if blockType >= wbtContext && blockType <= wbtExtension { + headerLen = 8 + } + + if blockLen < headerLen { + break + } + content := data[offset+headerLen : offset+blockLen] + + switch blockType { + case wbtSync, wbtCodecVersions, wbtChannels, wbtContext, + wbtFrameBegin, wbtFrameEnd: + // Infrastructure blocks — no action needed for decoding. + case wbtRegion: + rects = d.parseRegion(content, left, top) + case wbtExtension: + quants = d.decodeTileset(content, left, top, surfData, width, height) + } + + offset += blockLen + } + + // If no rects were parsed from REGION (e.g. numRects=0), generate one + // covering the entire surface per MS-RDPRFX 2.2.2.3.3. + if len(rects) == 0 && quants != nil { + rects = []rfxRect{{x: left, y: top, w: width - left, h: height - top}} + } + + return rects +} + +// parseRegion extracts rectangles from a WBT_REGION block. +// left/top are the WTS1 destination offsets applied to produce surface coordinates. +func (d *rfxDecoder) parseRegion(data []byte, left, top int) []rfxRect { + if len(data) < 7 { + return nil + } + + // regionFlags := data[0] + numRects := binary.LittleEndian.Uint16(data[1:]) + + if numRects == 0 { + return nil + } + + needed := 3 + int(numRects)*8 + 4 + if len(data) < needed { + return nil + } + + if cap(d.rectsBuf) >= int(numRects) { + d.rectsBuf = d.rectsBuf[:numRects] + } else { + d.rectsBuf = make([]rfxRect, numRects) + } + rects := d.rectsBuf + off := 3 + for i := range numRects { + rects[i] = rfxRect{ + x: left + int(binary.LittleEndian.Uint16(data[off:])), + y: top + int(binary.LittleEndian.Uint16(data[off+2:])), + w: int(binary.LittleEndian.Uint16(data[off+4:])), + h: int(binary.LittleEndian.Uint16(data[off+6:])), + } + off += 8 + } + + // Validate regionType + regionType := binary.LittleEndian.Uint16(data[off:]) + if regionType != cbtRegion { + slog.Debug("RFX: unexpected regionType", "type", regionType) + } + + return rects +} + +// decodeTileset parses and decodes all tiles from a WBT_EXTENSION/TILESET block. +// Format: subtype(2) + idx(2) + properties(2) + numQuant(1) + tileSize(1) + +// +// numTiles(2) + tilesDataSize(4) + quants(numQuant*5) + tiles +// +// Returns the quant table for caller reference. +func (d *rfxDecoder) decodeTileset(data []byte, left, top int, surfData []byte, width, height int) []rfxQuant { + if len(data) < 14 { + return nil + } + + subtype := binary.LittleEndian.Uint16(data[0:]) + if subtype != cbtTileset { + return nil + } + + properties := binary.LittleEndian.Uint16(data[4:]) + numQuant := int(data[6]) + // tileSize := data[7] + numTiles := int(binary.LittleEndian.Uint16(data[8:])) + // tilesDataSize := binary.LittleEndian.Uint32(data[10:]) + + // Extract RLGR entropy algorithm from TILESET properties. + // TILESET properties bit layout (MS-RDPRFX / FreeRDP): + // bits 10-13: et (entropy type) - 0x01=RLGR1, 0x04=RLGR3 + rlgrMode := 1 + et := (properties >> 10) & 0x0F + if et == 0x04 { + rlgrMode = 3 + } + + off := 14 + + // Parse quantization tables (5 bytes each, 10 nibbles) + if off+numQuant*5 > len(data) { + return nil + } + if cap(d.quantsBuf) >= numQuant { + d.quantsBuf = d.quantsBuf[:numQuant] + } else { + d.quantsBuf = make([]rfxQuant, numQuant) + } + quants := d.quantsBuf + for i := range numQuant { + quants[i] = parseRfxQuant(data[off:]) + off += 5 + } + + // Collect tile content slices for parallel decoding. + if cap(d.tilesBuf) >= numTiles { + d.tilesBuf = d.tilesBuf[:0] + } else { + d.tilesBuf = make([]rfxTileWork, 0, numTiles) + } + tiles := d.tilesBuf + for range numTiles { + if off+6 > len(data) { + break + } + tileBlockType := binary.LittleEndian.Uint16(data[off:]) + tileBlockLen := int(binary.LittleEndian.Uint32(data[off+2:])) + + if tileBlockType != cbtTile { + break + } + if tileBlockLen < 19 || off+tileBlockLen > len(data) { + break + } + + tiles = append(tiles, rfxTileWork{content: data[off+6 : off+tileBlockLen]}) + off += tileBlockLen + } + d.tilesBuf = tiles + + // Decode tiles concurrently — each tile writes to its own non-overlapping + // 64×64 region of the output buffer so no locking is needed. For small + // tile counts the goroutine + channel + WaitGroup overhead exceeds the + // per-tile work, so fall back to serial decoding below the threshold. + const parallelTileThreshold = 12 + if len(tiles) >= parallelTileThreshold { + workers := min(runtime.NumCPU(), len(tiles)) + ch := make(chan rfxTileWork, len(tiles)) + for _, t := range tiles { + ch <- t + } + close(ch) + var wg sync.WaitGroup + for range workers { + wg.Go(func() { + defer func() { + if r := recover(); r != nil { + slog.Error("RFX: tile decode panic", "err", r) + } + }() + for t := range ch { + d.decodeTile(t.content, quants, rlgrMode, left, top, surfData, width, height, false) + } + }) + } + wg.Wait() + } else { + for _, t := range tiles { + d.decodeTile(t.content, quants, rlgrMode, left, top, surfData, width, height, true) + } + } + + return quants +} + +// decodeTile decodes a single non-progressive RFX tile. +// Format: quantIdxY(1) + quantIdxCb(1) + quantIdxCr(1) + xIdx(2) + yIdx(2) + +// +// YLen(2) + CbLen(2) + CrLen(2) + YData(YLen) + CbData(CbLen) + CrData(CrLen) +// +// When parallelComponents is true the Y, Cb, and Cr channels are decoded +// concurrently (safe because each works on its own independent data and pool +// buffer). Use true for the serial-tile path; false when the outer worker pool +// already saturates all CPUs. +func (d *rfxDecoder) decodeTile(data []byte, quants []rfxQuant, rlgrMode int, left, top int, output []byte, outW, outH int, parallelComponents bool) { + if len(data) < 13 { + return + } + + quantIdxY := int(data[0]) + quantIdxCb := int(data[1]) + quantIdxCr := int(data[2]) + xIdx := int(binary.LittleEndian.Uint16(data[3:])) + yIdx := int(binary.LittleEndian.Uint16(data[5:])) + yLen := int(binary.LittleEndian.Uint16(data[7:])) + cbLen := int(binary.LittleEndian.Uint16(data[9:])) + crLen := int(binary.LittleEndian.Uint16(data[11:])) + + off := 13 + yData := safeSlice(data, off, yLen) + off += yLen + cbData := safeSlice(data, off, cbLen) + off += cbLen + crData := safeSlice(data, off, crLen) + + qY := rfxGetQuant(quants, quantIdxY) + qCb := rfxGetQuant(quants, quantIdxCb) + qCr := rfxGetQuant(quants, quantIdxCr) + + var yPixels, cbPixels, crPixels []int16 + if parallelComponents { + var wg sync.WaitGroup + wg.Go(func() { yPixels = rfxDecodeComponent(yData, qY, rlgrMode) }) + wg.Go(func() { cbPixels = rfxDecodeComponent(cbData, qCb, rlgrMode) }) + wg.Go(func() { crPixels = rfxDecodeComponent(crData, qCr, rlgrMode) }) + wg.Wait() + } else { + yPixels = rfxDecodeComponent(yData, qY, rlgrMode) + cbPixels = rfxDecodeComponent(cbData, qCb, rlgrMode) + crPixels = rfxDecodeComponent(crData, qCr, rlgrMode) + } + + // Apply WTS1 left/top offset: tile pixel position on surface = + // left + xIdx*64, top + yIdx*64 (per FreeRDP/MS-RDPRFX). + rfxPlaceTileAbs(yPixels, cbPixels, crPixels, left+xIdx*rfxTileSize, top+yIdx*rfxTileSize, output, outW, outH) + + coeffPool.Put((*coeffArr)(yPixels)) + coeffPool.Put((*coeffArr)(cbPixels)) + coeffPool.Put((*coeffArr)(crPixels)) +} + +// DecodeSurfaceRFX decodes non-progressive RemoteFX (MS-RDPRFX) encoded data +// into a top-down BGRA pixel buffer suitable for surface bitmap commands. +func DecodeSurfaceRFX(data []byte, width, height int) []byte { + output := make([]byte, width*height*4) + dec := newRfxDecoder() + dec.Decode(data, 0, 0, output, width, height) + return output +} diff --git a/plugin/rdpgfx/rfx_dwt_shared.go b/plugin/rdpgfx/rfx_dwt_shared.go new file mode 100644 index 0000000..a09c9c2 --- /dev/null +++ b/plugin/rdpgfx/rfx_dwt_shared.go @@ -0,0 +1,221 @@ +package rdpgfx + +// Standard RemoteFX (MS-RDPRFX) codec helpers shared with rfx.go. +// Extracted from the progressive decoder rewrite; algorithms unchanged. + +func rfxGetQuant(quants []rfxQuant, idx int) rfxQuant { + if idx < len(quants) { + return quants[idx] + } + return rfxQuant{6, 6, 6, 6, 6, 6, 6, 6, 6, 6} +} +// rfxDecodeComponent decodes one color component (Y, Cb, or Cr) for a 64×64 tile. +// The returned slice is backed by a *coeffArr from coeffPool; the caller must +// return it via coeffPool.Put((*coeffArr)(result)) when done. +func rfxDecodeComponent(data []byte, quant rfxQuant, rlgrMode int) []int16 { + const tilePixels = rfxTileSize * rfxTileSize // 4096 + + // Get a pooled coefficient buffer. The pool stores *coeffArr (pointer to a + // fixed-size array) so the any interface stores a single pointer word with no + // heap-boxing allocation. + arr := coeffPool.Get().(*coeffArr) + coeffs := arr[:] + + if data == nil { + clear(coeffs) + return coeffs + } + + // 1. RLGR entropy decode → 4096 coefficients + if rlgrMode == 3 { + coeffs = rlgr3Decode(data, tilePixels, coeffs) + } else { + coeffs = rlgr1Decode(data, tilePixels, coeffs) + } + + // 2. Differential decode LL3 and dequantize LL3 in a single pass. + // Mathematical identity: cumsum(x) * 2^s == cumsum_of(x * 2^s) + // so we can left-shift each element before accumulating. + if quant.LL3 > 1 { + shift := quant.LL3 - 1 + coeffs[4032] <<= shift + for i := 4033; i < 4096; i++ { + coeffs[i] = coeffs[i-1] + coeffs[i]<>1) + prevEvenH := lh[rowOff] - int16((int32(hh[rowOff])*2+1)>>1) + tmp[lDstOff] = prevEvenL + tmp[hDstOff] = prevEvenH + + // col=1..n-1: compute even[col], then immediately compute odd[col-1] + // using prevEven (=even[col-1], still in register) and the just-computed + // even[col] — no re-read of tmp required. + for col := 1; col < n; col++ { + x := col << 1 + evenL := ll[rowOff+col] - int16((int32(hl[rowOff+col-1])+int32(hl[rowOff+col])+1)>>1) + evenH := lh[rowOff+col] - int16((int32(hh[rowOff+col-1])+int32(hh[rowOff+col])+1)>>1) + tmp[lDstOff+x-1] = int16((int32(hl[rowOff+col-1])<<1) + ((int32(prevEvenL)+int32(evenL))>>1)) + tmp[hDstOff+x-1] = int16((int32(hh[rowOff+col-1])<<1) + ((int32(prevEvenH)+int32(evenH))>>1)) + tmp[lDstOff+x] = evenL + tmp[hDstOff+x] = evenH + prevEvenL = evenL + prevEvenH = evenH + } + + // last odd[n-1]: right boundary, even[n] = even[n-1]. + x := (n - 1) << 1 + tmp[lDstOff+x+1] = int16((int32(hl[rowOff+n-1])<<1) + int32(prevEvenL)) + tmp[hDstOff+x+1] = int16((int32(hh[rowOff+n-1])<<1) + int32(prevEvenH)) + } + + // Step 2: Vertical IDWT on each column. + // Process 8 columns at a time to improve cache utilisation — a cache line + // holds 32 int16 values; 8 columns keeps the working set within one or two + // lines per row access. All valid sizes (16, 32, 64) divide evenly by 8, + // so the scalar tail loop is never reached in practice. + const blk = 8 + col := 0 + for ; col+blk <= size; col += blk { + // Row 0: first even output (no previous odd) + l0 := tmp[col : col+blk] + h0 := tmp[n*size+col : n*size+col+blk] + out0 := buf[col : col+blk] + for b := range blk { + out0[b] = int16(int32(l0[b]) - ((int32(h0[b])*2 + 1) >> 1)) + } + // Rows 1..n-1: interleaved even/odd outputs + for row := 1; row < n; row++ { + lBase := row*size + col + hBase := (row+n)*size + col + hPrevBase := (row-1+n)*size + col + evenBase := 2*row*size + col + prevEvenBase := (2*row-2)*size + col + oddBase := (2*row-1)*size + col + + l := tmp[lBase : lBase+blk] + h := tmp[hBase : hBase+blk] + hPrev := tmp[hPrevBase : hPrevBase+blk] + evenOut := buf[evenBase : evenBase+blk] + prevEvenIn := buf[prevEvenBase : prevEvenBase+blk] + oddOut := buf[oddBase : oddBase+blk] + + for b := range blk { + hPrevV := int32(hPrev[b]) + even := int32(l[b]) - ((hPrevV + int32(h[b]) + 1) >> 1) + evenOut[b] = int16(even) + oddOut[b] = int16((hPrevV << 1) + ((int32(prevEvenIn[b]) + even) >> 1)) + } + } + // Last odd row + lastEvenBase := (2*n-2)*size + col + lastHBase := (2*n-1)*size + col + lastEvenSlice := buf[lastEvenBase : lastEvenBase+blk] + lastHSlice := tmp[lastHBase : lastHBase+blk] + lastOddOut := buf[lastHBase : lastHBase+blk] + for b := range blk { + lastOddOut[b] = int16((int32(lastHSlice[b]) << 1) + int32(lastEvenSlice[b])) + } + } + for ; col < size; col++ { + lVal := int32(tmp[col]) + hVal := int32(tmp[n*size+col]) + buf[col] = int16(lVal - ((hVal*2 + 1) >> 1)) + + for row := 1; row < n; row++ { + lIdx := row*size + col + hIdx := (row+n)*size + col + hPrevIdx := (row-1+n)*size + col + + even := int32(tmp[lIdx]) - ((int32(tmp[hPrevIdx]) + int32(tmp[hIdx]) + 1) >> 1) + buf[2*row*size+col] = int16(even) + + prevEven := int32(buf[(2*row-2)*size+col]) + odd := (int32(tmp[hPrevIdx]) << 1) + ((prevEven + even) >> 1) + buf[(2*row-1)*size+col] = int16(odd) + } + + lastEven := int32(buf[(2*n-2)*size+col]) + lastH := int32(tmp[(2*n-1)*size+col]) + buf[(2*n-1)*size+col] = int16((lastH << 1) + lastEven) + } +} + +func rfxShiftSubband(data []int16, factor uint8) { + if factor <= 1 { + return + } + shift := factor - 1 + for i := range data { + data[i] <<= shift + } +} + diff --git a/plugin/rdpgfx/rfx_pool.go b/plugin/rdpgfx/rfx_pool.go new file mode 100644 index 0000000..68c6060 --- /dev/null +++ b/plugin/rdpgfx/rfx_pool.go @@ -0,0 +1,50 @@ +package rdpgfx + +import "sync" + +// Buffer pools for RFX tile decoding to minimize allocations in the hot path. +// Each 64×64 tile needs 4096 int16 coefficients per component (Y, Cb, Cr) +// and several temporary buffers for the IDWT. + +// coeffArr is the fixed-size coefficient array stored in the pool. +// Using a pointer to an array (*coeffArr) avoids interface-boxing allocations: +// a pointer fits in one word of the any interface, whereas a []int16 header +// (pointer + len + cap = 24 bytes) always requires a 24-byte heap box. +type coeffArr = [4096]int16 + +var coeffPool = sync.Pool{ + New: func() any { return new(coeffArr) }, +} + +// idwtBufs holds the temporary buffer for one rfxIDWT2DLevel call. +// The subbands (HL/LH/HH/LL) are read directly from the input buf without copying. +type idwtBufs struct { + tmp []int16 // intermediate row-interleaved buffer; max 64×64 = 4096 +} + +var idwtBufPool = sync.Pool{ + New: func() any { + return &idwtBufs{tmp: make([]int16, 4096)} + }, +} + +// planeBufPool reuses byte slices for the per-component planes in decodePlanar. +// The plane size varies (w×h bytes), so the pool stores capacity-keyed slices; +// acquirePlaneBuf re-slices to exactly `size` when the capacity is sufficient. +var planeBufPool = sync.Pool{ + New: func() any { return []byte(nil) }, +} + +func acquirePlaneBuf(size int) []byte { + b := planeBufPool.Get().([]byte) + if cap(b) >= size { + return b[:size] + } + return make([]byte, size) +} + +func releasePlaneBuf(b []byte) { + if b != nil { + planeBufPool.Put(b[:cap(b)]) + } +} diff --git a/plugin/rdpgfx/rfx_progressive.go b/plugin/rdpgfx/rfx_progressive.go new file mode 100644 index 0000000..008bc93 --- /dev/null +++ b/plugin/rdpgfx/rfx_progressive.go @@ -0,0 +1,1321 @@ +package rdpgfx + +// RFX Progressive Codec decoder (MS-RDPEGFX 2.2.4), algorithm aligned with +// FreeRDP libfreerdp/codec/progressive.c. Handles RDPGFX_CODECID_CAPROGRESSIVE +// (0x0009) in WIRE_TO_SURFACE_PDU_1/2. +// +// Key points that differ from the plain RemoteFX codec (MS-RDPRFX): +// - The 5-byte quant table uses the RDPEGFX band order (LL3,HL3,LH3,HH3, +// HL2,LH2,HH2,HL1,LH1,HH1), which swaps LH3/HL3, HL2/LH2 and HL1/LH1 +// compared to the RDPRFX order parsed by parseRfxQuant. +// - The effective dequant shift per band is (plain + progressive quant) - 1. +// - TILE_FIRST/TILE_SIMPLE carry the first pass of a tile; TILE_UPGRADE +// carries incremental bit-plane refinements using an SRL+RAW dual bit +// stream applied to the cached coefficients in extrapolate layout. +// - Region flag RFX_DWT_REDUCE_EXTRAPOLATE switches the IDWT to the +// extrapolate variant with irregular band sizes. + +import ( + "encoding/binary" + "fmt" + "log/slog" + "runtime" + "sync" +) + +// Progressive block types (different from non-progressive WBT_* at same values!) +const ( + progWBTSync = 0xCCC0 + progWBTFrameBegin = 0xCCC1 + progWBTFrameEnd = 0xCCC2 + progWBTContext = 0xCCC3 + progWBTRegion = 0xCCC4 + progWBTTileSimple = 0xCCC5 + progWBTTileFirst = 0xCCC6 + progWBTTileUpgrade = 0xCCC7 +) + +// Region/context flags (MS-RDPEGFX progressive.h) +const ( + progFlagSubbandDiffing = 0x01 // PROGRESSIVE_BLOCK_CONTEXT::flags + progFlagDWTReduceExtrapolate = 0x01 // PROGRESSIVE_BLOCK_REGION::flags + progFlagTileDifference = 0x01 // tile flags +) + +const rfxTileSize = 64 + +// rfxQuant holds the 10 quantization values in RDPRFX order (standard codec). +type rfxQuant struct { + LL3, LH3, HL3, HH3 uint8 + LH2, HL2, HH2 uint8 + LH1, HL1, HH1 uint8 +} + +// parseRfxQuant parses a standard (MS-RDPRFX) 5-byte quant table. +// Used by the plain RemoteFX tileset decoder in rfx.go. +func parseRfxQuant(data []byte) rfxQuant { + return rfxQuant{ + LL3: data[0] & 0x0F, + LH3: data[0] >> 4, + HL3: data[1] & 0x0F, + HH3: data[1] >> 4, + LH2: data[2] & 0x0F, + HL2: data[2] >> 4, + HH2: data[3] & 0x0F, + LH1: data[3] >> 4, + HL1: data[4] & 0x0F, + HH1: data[4] >> 4, + } +} + +// progBandQuant holds the 10 band quant values in RDPEGFX order +// (RFX_COMPONENT_CODEC_QUANT in FreeRDP progressive.h). +type progBandQuant struct { + LL3, HL3, LH3, HH3 uint8 + HL2, LH2, HH2 uint8 + HL1, LH1, HH1 uint8 +} + +// parseProgBandQuant reads one 5-byte progressive quant component. +func parseProgBandQuant(data []byte) progBandQuant { + return progBandQuant{ + LL3: data[0] & 0x0F, + HL3: data[0] >> 4, + LH3: data[1] & 0x0F, + HH3: data[1] >> 4, + HL2: data[2] & 0x0F, + LH2: data[2] >> 4, + HH2: data[3] & 0x0F, + HL1: data[3] >> 4, + LH1: data[4] & 0x0F, + HH1: data[4] >> 4, + } +} + +// progCodecQuant is one 16-byte RFX_PROGRESSIVE_CODEC_QUANT entry. +type progCodecQuant struct { + quality byte + y, cb, cr progBandQuant +} + +// progAdd returns a+b band-wise. +func progAdd(a, b progBandQuant) progBandQuant { + return progBandQuant{ + a.LL3 + b.LL3, a.HL3 + b.HL3, a.LH3 + b.LH3, a.HH3 + b.HH3, + a.HL2 + b.HL2, a.LH2 + b.LH2, a.HH2 + b.HH2, + a.HL1 + b.HL1, a.LH1 + b.LH1, a.HH1 + b.HH1, + } +} + +// progSub returns a-b band-wise, ok=false when any band would underflow. +func progSub(a, b progBandQuant) (progBandQuant, bool) { + if a.LL3 < b.LL3 || a.HL3 < b.HL3 || a.LH3 < b.LH3 || a.HH3 < b.HH3 || + a.HL2 < b.HL2 || a.LH2 < b.LH2 || a.HH2 < b.HH2 || + a.HL1 < b.HL1 || a.LH1 < b.LH1 || a.HH1 < b.HH1 { + return progBandQuant{}, false + } + return progBandQuant{ + a.LL3 - b.LL3, a.HL3 - b.HL3, a.LH3 - b.LH3, a.HH3 - b.HH3, + a.HL2 - b.HL2, a.LH2 - b.LH2, a.HH2 - b.HH2, + a.HL1 - b.HL1, a.LH1 - b.LH1, a.HH1 - b.HH1, + }, true +} + +// progLSub subtracts v from every band, ok=false on underflow or out-of-range v. +func progLSub(a progBandQuant, v int) (progBandQuant, bool) { + if v < 0 || v > 255 { + return progBandQuant{}, false + } + return progSub(a, progBandQuant{ + LL3: uint8(v), HL3: uint8(v), LH3: uint8(v), HH3: uint8(v), + HL2: uint8(v), LH2: uint8(v), HH2: uint8(v), + HL1: uint8(v), LH1: uint8(v), HH1: uint8(v), + }) +} + +// progTileState caches the progressive state of one tile (per component: +// current coefficients and first-pass sign array). +type progTileState struct { + pass int + // per component: 0=Y, 1=Cb, 2=Cr + current [3]*coeffArr // accumulated band coefficients (extrapolate layout) + sign [3]*coeffArr // raw first-pass RLGR output (signs for upgrades) + yBitPos progBandQuant + cbBitPos progBandQuant + crBitPos progBandQuant +} + +type rfxTileCoeffs = progTileState + +type rfxProgTileWork struct { + tileType uint16 + data []byte +} + +type progRegionCtx struct { + quantVals []progBandQuant // numQuant entries (plain quants) + quantProgVals []progCodecQuant // numProgQuant entries + numQuant int + numProgQuant int + flags byte // RFX_DWT_REDUCE_EXTRAPOLATE + extrapolate bool + rects []rfxRect // 脏矩形:瓦片渲染必须裁剪到其并集内 +} + +type rfxProgressiveDecoder struct { + mu sync.RWMutex + tileCache map[uint32]*progTileState // key: yIdx<<16 | xIdx + rectsBuf []rfxRect + quantsBuf []progBandQuant + progQuantsBuf []progCodecQuant + tilesBuf []rfxProgTileWork + contextFlags byte // PROGRESSIVE_BLOCK_CONTEXT flags + logged [32]bool +} + +// logOnce 每类失败只打第一条日志(复用 clearCodecCtx 的槽位思路) +func (d *rfxProgressiveDecoder) logOnce(slot int, msg string, args ...any) { + d.mu.Lock() + defer d.mu.Unlock() + if d.logged[slot] { + return + } + d.logged[slot] = true + slog.Warn("progressive:"+msg, args...) +} + +func newRfxProgressiveDecoder() *rfxProgressiveDecoder { + return &rfxProgressiveDecoder{ + tileCache: make(map[uint32]*progTileState), + } +} + +// Reset discards the tile coefficient cache. Call this whenever the server +// starts a new progressive sequence (e.g. on RESET_GRAPHICS). +func (d *rfxProgressiveDecoder) Reset() { + d.mu.Lock() + old := d.tileCache + d.tileCache = make(map[uint32]*progTileState) + d.mu.Unlock() + for _, ts := range old { + progFreeTileState(ts) + } +} + +func progFreeTileState(ts *progTileState) { + if ts == nil { + return + } + for c := 0; c < 3; c++ { + if ts.current[c] != nil { + coeffPool.Put(ts.current[c]) + } + if ts.sign[c] != nil { + coeffPool.Put(ts.sign[c]) + } + } +} + +// rfxRect represents a rectangle of decoded tiles. +type rfxRect struct { + x, y, w, h int +} + +// Decode processes RFX Progressive codec data, rendering tiles onto the +// provided surface buffer. Returns the bounding rectangles of decoded regions. +func (d *rfxProgressiveDecoder) Decode(data []byte, surfData []byte, width, height int) []rfxRect { + var rects []rfxRect + + offset := 0 + for offset+6 <= len(data) { + blockType := binary.LittleEndian.Uint16(data[offset:]) + blockLen := binary.LittleEndian.Uint32(data[offset+2:]) + + if blockLen < 6 || offset+int(blockLen) > len(data) { + break + } + + blockData := data[offset+6 : offset+int(blockLen)] + + switch blockType { + case progWBTSync: + // magic + version — nothing to do. + case progWBTFrameBegin, progWBTFrameEnd: + // frame bookkeeping — nothing to do. + case progWBTContext: + // ctxId(1) + tileSize(2) + flags(1) + if len(blockData) >= 4 { + d.contextFlags = blockData[3] + } + case progWBTRegion: + regionRects, _ := d.parseRegion(blockData, surfData, width, height) + rects = append(rects, regionRects...) + default: + slog.Debug("RFX: unknown progressive block type", "type", blockType) + } + + offset += int(blockLen) + } + + return rects +} + +// parseRegion extracts rects and quant tables from a PROGRESSIVE_WBT_REGION block, +// and decodes the tile sub-blocks embedded within it onto the surface. +func (d *rfxProgressiveDecoder) parseRegion(data []byte, surfData []byte, outW, outH int) ([]rfxRect, []progBandQuant) { + if len(data) < 12 { + return nil, nil + } + + // tileSize := data[0] + numRects := int(binary.LittleEndian.Uint16(data[1:])) + numQuant := int(data[3]) + numProgQuant := int(data[4]) + flags := data[5] + numTiles := int(binary.LittleEndian.Uint16(data[6:])) + // tileDataSize := binary.LittleEndian.Uint32(data[8:]) + + offset := 12 + extrapolate := flags&progFlagDWTReduceExtrapolate != 0 + region := progRegionCtx{ + numQuant: numQuant, + numProgQuant: numProgQuant, + flags: flags, + extrapolate: extrapolate, + } + + // Parse rects (8 bytes each: x, y, width, height as uint16) + if cap(d.rectsBuf) >= numRects { + d.rectsBuf = d.rectsBuf[:numRects] + } else { + d.rectsBuf = make([]rfxRect, numRects) + } + rects := d.rectsBuf + for i := range numRects { + if offset+8 > len(data) { + return nil, nil + } + rx := int(binary.LittleEndian.Uint16(data[offset:])) + ry := int(binary.LittleEndian.Uint16(data[offset+2:])) + rw := int(binary.LittleEndian.Uint16(data[offset+4:])) + rh := int(binary.LittleEndian.Uint16(data[offset+6:])) + rects[i] = rfxRect{x: rx, y: ry, w: rw, h: rh} + offset += 8 + } + region.rects = rects + + // Parse plain quant values (5 bytes each, RDPEGFX band order) + if cap(d.quantsBuf) >= numQuant { + d.quantsBuf = d.quantsBuf[:numQuant] + } else { + d.quantsBuf = make([]progBandQuant, numQuant) + } + quants := d.quantsBuf + for i := range numQuant { + if offset+5 > len(data) { + return nil, nil + } + quants[i] = parseProgBandQuant(data[offset:]) + offset += 5 + } + region.quantVals = quants + + // Parse progressive quant values (16 bytes each: quality + 3 components) + if cap(d.progQuantsBuf) >= numProgQuant { + d.progQuantsBuf = d.progQuantsBuf[:numProgQuant] + } else { + d.progQuantsBuf = make([]progCodecQuant, numProgQuant) + } + progQuants := d.progQuantsBuf + for i := range numProgQuant { + if offset+16 > len(data) { + return nil, nil + } + progQuants[i].quality = data[offset] + progQuants[i].y = parseProgBandQuant(data[offset+1:]) + progQuants[i].cb = parseProgBandQuant(data[offset+6:]) + progQuants[i].cr = parseProgBandQuant(data[offset+11:]) + offset += 16 + } + region.quantProgVals = progQuants + + // Collect all decodable tiles before dispatching, so we can parallelise + // when there are enough to amortise goroutine overhead (same threshold as + // non-progressive decodeTileset in rfx.go). + if cap(d.tilesBuf) >= numTiles { + d.tilesBuf = d.tilesBuf[:0] + } else { + d.tilesBuf = make([]rfxProgTileWork, 0, numTiles) + } + tiles := d.tilesBuf + for offset+6 <= len(data) { + tileType := binary.LittleEndian.Uint16(data[offset:]) + tileLen := binary.LittleEndian.Uint32(data[offset+2:]) + if tileLen < 6 || offset+int(tileLen) > len(data) { + break + } + switch tileType { + case progWBTTileSimple, progWBTTileFirst, progWBTTileUpgrade: + tiles = append(tiles, rfxProgTileWork{tileType: tileType, data: data[offset+6 : offset+int(tileLen)]}) + default: + slog.Debug("RFX: unknown progressive tile type", "type", tileType) + } + offset += int(tileLen) + } + d.tilesBuf = tiles + + const parallelTileThreshold = 12 + decodeTile := func(tw rfxProgTileWork, parallel bool) { + switch tw.tileType { + case progWBTTileSimple: + d.decodeTileSimple(tw.data, ®ion, surfData, outW, outH, parallel) + case progWBTTileFirst: + d.decodeTileFirst(tw.data, ®ion, surfData, outW, outH, parallel) + case progWBTTileUpgrade: + d.decodeTileUpgrade(tw.data, ®ion, surfData, outW, outH, parallel) + } + } + if len(tiles) >= parallelTileThreshold { + workers := min(runtime.NumCPU(), len(tiles)) + ch := make(chan rfxProgTileWork, len(tiles)) + for _, tw := range tiles { + ch <- tw + } + close(ch) + var wg sync.WaitGroup + for range workers { + wg.Go(func() { + defer func() { + if r := recover(); r != nil { + slog.Error("RFX progressive: tile decode panic", "err", r) + } + }() + for tw := range ch { + decodeTile(tw, false) + } + }) + } + wg.Wait() + } else { + for _, tw := range tiles { + decodeTile(tw, true) + } + } + + return rects, quants +} + +// safeSlice returns data[offset:offset+length] when fully in range, else nil. +// Shared with the standard RemoteFX tileset decoder in rfx.go. +func safeSlice(data []byte, offset, length int) []byte { + if length <= 0 || offset < 0 || offset+length > len(data) { + return nil + } + return data[offset : offset+length] +} + +// progTileHeader is the common prefix of all tile block headers. +type progTileHeader struct { + quantIdxY byte + quantIdxCb byte + quantIdxCr byte + xIdx int + yIdx int + flags byte + quality byte // 0xFF = full quality (quantProgValFull) + // simple/first + yLen, cbLen, crLen, tailLen int + // upgrade + ySrlLen, yRawLen int + cbSrlLen, cbRawLen int + crSrlLen, crRawLen int +} + +// getProgTileState returns the cache entry for the tile, allocating a fresh +// state (releasing the old one) for a new SIMPLE/FIRST pass. +func (d *rfxProgressiveDecoder) getProgTileState(key uint32, firstPass bool) *progTileState { + d.mu.Lock() + defer d.mu.Unlock() + // 注意:FIRST pass 到达时不得销毁已有 state——服务器可能对同一瓦片 + // 连续发送多个 FIRST(如 RFX_TILE_DIFFERENCE),其差分系数基于客户端 + // 应持有的参考状态(Ref);重置会导致差分失去基准而产生花屏块。 + ts := d.tileCache[key] + _ = firstPass + if ts == nil { + ts = &progTileState{} + for c := 0; c < 3; c++ { + ts.current[c] = coeffPool.Get().(*coeffArr) + ts.sign[c] = coeffPool.Get().(*coeffArr) + clear(ts.current[c][:]) + clear(ts.sign[c][:]) + } + ts.pass = 0 + ts.yBitPos, ts.cbBitPos, ts.crBitPos = progBandQuant{}, progBandQuant{}, progBandQuant{} + d.tileCache[key] = ts + } + return ts +} + +// selectQuant resolves the plain and progressive quant sets for a +// SIMPLE/FIRST tile. FIRST passes never consult the previous tile state: +// original tiles replace the reference and difference tiles add to it, so +// the numBits bookkeeping (an UPGRADE-only concept, FreeRDP computes it in +// progressive_rfx_upgrade_component only) must not gate FIRST decoding. +func selectQuant(region *progRegionCtx, hdr *progTileHeader, + which int) (plain, prog, shift progBandQuant, ok bool) { + var idx byte + switch which { + case 0: + idx = hdr.quantIdxY + case 1: + idx = hdr.quantIdxCb + default: + idx = hdr.quantIdxCr + } + if int(idx) >= len(region.quantVals) { + return progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + plain = region.quantVals[idx] + if hdr.quality == 0xFF { + prog = progBandQuant{} // quantProgValFull is all-zero in FreeRDP + } else { + if int(hdr.quality) >= len(region.quantProgVals) { + return progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + switch which { + case 0: + prog = region.quantProgVals[hdr.quality].y + case 1: + prog = region.quantProgVals[hdr.quality].cb + default: + prog = region.quantProgVals[hdr.quality].cr + } + } + combined := progAdd(plain, prog) + // 量化值为 0 的波段按规范不编码(MS-RDPRFX),该带 shift 不会应用到任何 + // 系数;钳到 0 兼容组合值为 0 的区域量化,而不是丢弃整个瓦片。 + shift = progShiftClamped(combined, 1) + return plain, prog, shift, true +} + +// progShiftClamped returns max(a-v, 0) band-wise. +func progShiftClamped(a progBandQuant, v int) progBandQuant { + sub := func(x uint8) uint8 { + s := int(x) - v + if s < 0 { + return 0 + } + return uint8(s) + } + return progBandQuant{ + LL3: sub(a.LL3), HL3: sub(a.HL3), LH3: sub(a.LH3), HH3: sub(a.HH3), + HL2: sub(a.HL2), LH2: sub(a.LH2), HH2: sub(a.HH2), + HL1: sub(a.HL1), LH1: sub(a.LH1), HH1: sub(a.HH1), + } +} + +// decodeTileSimple handles PROGRESSIVE_WBT_TILE_SIMPLE (0xCCC5). +func (d *rfxProgressiveDecoder) decodeTileSimple(data []byte, region *progRegionCtx, output []byte, outW, outH int, parallelComponents bool) { + d.decodeTileFirstPass(data, region, output, outW, outH, parallelComponents, progWBTTileSimple) +} + +// decodeTileFirst handles PROGRESSIVE_WBT_TILE_FIRST (0xCCC6). +func (d *rfxProgressiveDecoder) decodeTileFirst(data []byte, region *progRegionCtx, output []byte, outW, outH int, parallelComponents bool) { + d.decodeTileFirstPass(data, region, output, outW, outH, parallelComponents, progWBTTileFirst) +} + +// decodeTileFirstPass implements the shared SIMPLE/FIRST logic (first pass of +// a tile progression). +func (d *rfxProgressiveDecoder) decodeTileFirstPass(data []byte, region *progRegionCtx, output []byte, outW, outH int, parallelComponents bool, tileType uint16) { + hdrLen := 16 + if tileType == progWBTTileFirst { + hdrLen = 17 + } + if len(data) < hdrLen+1 { + return + } + hdr := progTileHeader{ + quantIdxY: data[0], + quantIdxCb: data[1], + quantIdxCr: data[2], + xIdx: int(binary.LittleEndian.Uint16(data[3:])), + yIdx: int(binary.LittleEndian.Uint16(data[5:])), + flags: data[7], + quality: 0xFF, + } + // Length fields precede the payload: the SIMPLE header is 16 bytes + // (yLen@8, cbLen@10, crLen@12, tailLen@14), FIRST inserts quality@8 and + // shifts the lengths to offsets 9/11/13/15 with a 17-byte header. + if tileType == progWBTTileFirst { + hdr.quality = data[8] + hdr.yLen = int(binary.LittleEndian.Uint16(data[9:])) + hdr.cbLen = int(binary.LittleEndian.Uint16(data[11:])) + hdr.crLen = int(binary.LittleEndian.Uint16(data[13:])) + hdr.tailLen = int(binary.LittleEndian.Uint16(data[15:])) + } else { + hdr.yLen = int(binary.LittleEndian.Uint16(data[8:])) + hdr.cbLen = int(binary.LittleEndian.Uint16(data[10:])) + hdr.crLen = int(binary.LittleEndian.Uint16(data[12:])) + hdr.tailLen = int(binary.LittleEndian.Uint16(data[14:])) + } + off := hdrLen + + yData := safeSlice(data, off, hdr.yLen) + off += hdr.yLen + cbData := safeSlice(data, off, hdr.cbLen) + off += hdr.cbLen + crData := safeSlice(data, off, hdr.crLen) + + key := uint32(hdr.yIdx)<<16 | uint32(hdr.xIdx) + ts := d.getProgTileState(key, true) + + var shifts, combineds [3]progBandQuant + for c := 0; c < 3; c++ { + plain, prog, shift, ok := selectQuant(region, &hdr, c) + if !ok { + return + } + shifts[c] = shift + combineds[c] = progAdd(plain, prog) + } + _ = combineds + + coeffDiff := hdr.flags&progFlagTileDifference != 0 + work := coeffPool.Get().(*coeffArr) + defer coeffPool.Put(work) + + // Decode each component; the DWT output lands in `work`, which we snapshot + // per component before the next component reuses the buffer. + var spatial [3]*coeffArr + spatial[0] = coeffPool.Get().(*coeffArr) + spatial[1] = coeffPool.Get().(*coeffArr) + spatial[2] = coeffPool.Get().(*coeffArr) + for c := 0; c < 3; c++ { + var compData []byte + switch c { + case 0: + compData = yData + case 1: + compData = cbData + default: + compData = crData + } + progDecodeComponent(compData, shifts[c], work, ts.sign[c], ts.current[c], coeffDiff, region.extrapolate) + copy(spatial[c][:], work[:]) + } + + rfxPlaceTile(spatial[0][:], spatial[1][:], spatial[2][:], hdr.xIdx, hdr.yIdx, output, outW, outH, region.rects) + + coeffPool.Put(spatial[0]) + coeffPool.Put(spatial[1]) + coeffPool.Put(spatial[2]) + + // FreeRDP: 每个 FIRST pass(含 DIFFERENCE)都把 pass 重置为 1,并把 + // bitPos 记为 quant+quantProg(组合位位置);UPGRADE 用 bitPos 差计算 + // numBits。 + ts.pass = 1 + ts.yBitPos = combineds[0] + ts.cbBitPos = combineds[1] + ts.crBitPos = combineds[2] +} + +func arrMin(a []int16) int16 { + m := a[0] + for _, v := range a { + if v < m { + m = v + } + } + return m +} + +func arrMax(a []int16) int16 { + m := a[0] + for _, v := range a { + if v > m { + m = v + } + } + return m +} + +// decodeTileUpgrade handles PROGRESSIVE_WBT_TILE_UPGRADE (0xCCC7): a 20-byte +// header followed by SRL/RAW stream pairs per component. +func (d *rfxProgressiveDecoder) decodeTileUpgrade(data []byte, region *progRegionCtx, output []byte, outW, outH int, parallelComponents bool) { + const hdrLen = 20 + if len(data) < hdrLen+1 { + return + } + hdr := progTileHeader{ + quantIdxY: data[0], + quantIdxCb: data[1], + quantIdxCr: data[2], + xIdx: int(binary.LittleEndian.Uint16(data[3:])), + yIdx: int(binary.LittleEndian.Uint16(data[5:])), + quality: data[7], + } + hdr.ySrlLen = int(binary.LittleEndian.Uint16(data[8:])) + hdr.yRawLen = int(binary.LittleEndian.Uint16(data[10:])) + hdr.cbSrlLen = int(binary.LittleEndian.Uint16(data[12:])) + hdr.cbRawLen = int(binary.LittleEndian.Uint16(data[14:])) + hdr.crSrlLen = int(binary.LittleEndian.Uint16(data[16:])) + hdr.crRawLen = int(binary.LittleEndian.Uint16(data[18:])) + + off := hdrLen + ySrl := safeSlice(data, off, hdr.ySrlLen) + off += hdr.ySrlLen + yRaw := safeSlice(data, off, hdr.yRawLen) + off += hdr.yRawLen + cbSrl := safeSlice(data, off, hdr.cbSrlLen) + off += hdr.cbSrlLen + cbRaw := safeSlice(data, off, hdr.cbRawLen) + off += hdr.cbRawLen + crSrl := safeSlice(data, off, hdr.crSrlLen) + off += hdr.crSrlLen + crRaw := safeSlice(data, off, hdr.crRawLen) + + key := uint32(hdr.yIdx)<<16 | uint32(hdr.xIdx) + ts := d.getProgTileState(key, false) + if ts.pass == 0 { + // Upgrade for a tile we never saw the first pass of: nothing to + // refine — skip rather than corrupt the cache. + d.logOnce(20, "upgrade skipped: no first pass", "x", hdr.xIdx, "y", hdr.yIdx) + return + } + + var shifts, numBitss, combineds [3]progBandQuant + for c := 0; c < 3; c++ { + plain, prog, shift, numBits, ok := selectQuantUpgrade(region, &hdr, ts, c) + if !ok { + d.logOnce(21, "upgrade quant resolve failed", "c", c, "quality", hdr.quality, + "quantIdxY", hdr.quantIdxY, "quantIdxCb", hdr.quantIdxCb, "quantIdxCr", hdr.quantIdxCr, + "nQuantVals", len(region.quantVals), "nQuantProgVals", len(region.quantProgVals), + "x", hdr.xIdx, "y", hdr.yIdx) + // 状态与 upgrade 目标不一致(此前的 pass 被丢弃或解析失败)。 + // 重置该瓦片,让下一个 FIRST 以全新基准重建,避免永久陈旧内容。 + ts.pass = 0 + ts.yBitPos, ts.cbBitPos, ts.crBitPos = progBandQuant{}, progBandQuant{}, progBandQuant{} + return + } + shifts[c] = shift + numBitss[c] = numBits + combineds[c] = progAdd(plain, prog) + } + + work := coeffPool.Get().(*coeffArr) + defer coeffPool.Put(work) + + var spatial [3]*coeffArr + spatial[0] = coeffPool.Get().(*coeffArr) + spatial[1] = coeffPool.Get().(*coeffArr) + spatial[2] = coeffPool.Get().(*coeffArr) + + for c := 0; c < 3; c++ { + var srlData, rawData []byte + switch c { + case 0: + srlData, rawData = ySrl, yRaw + case 1: + srlData, rawData = cbSrl, cbRaw + default: + srlData, rawData = crSrl, crRaw + } + progUpgradeComponent(work, ts.current[c], ts.sign[c], shifts[c], numBitss[c], srlData, rawData, region.extrapolate) + copy(spatial[c][:], work[:]) + } + + rfxPlaceTile(spatial[0][:], spatial[1][:], spatial[2][:], hdr.xIdx, hdr.yIdx, output, outW, outH, region.rects) + coeffPool.Put(spatial[0]) + coeffPool.Put(spatial[1]) + coeffPool.Put(spatial[2]) + + // 与 FIRST pass 相同:bitPos 记录组合位位置,供后续 UPGRADE 差分。 + ts.yBitPos = combineds[0] + ts.cbBitPos = combineds[1] + ts.crBitPos = combineds[2] + ts.pass++ +} + +// selectQuantUpgrade resolves shift/numBits for an upgrade pass. numBits = +// previous bit position - new combined bit position (the newly significant +// bits delivered by the upgrade stream). +func selectQuantUpgrade(region *progRegionCtx, hdr *progTileHeader, ts *progTileState, + which int) (progBandQuant, progBandQuant, progBandQuant, progBandQuant, bool) { + var idx byte + switch which { + case 0: + idx = hdr.quantIdxY + case 1: + idx = hdr.quantIdxCb + default: + idx = hdr.quantIdxCr + } + if int(idx) >= len(region.quantVals) { + return progBandQuant{}, progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + plain := region.quantVals[idx] + var prog progBandQuant + if hdr.quality == 0xFF { + prog = progBandQuant{} + } else { + if int(hdr.quality) >= len(region.quantProgVals) { + return progBandQuant{}, progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + switch which { + case 0: + prog = region.quantProgVals[hdr.quality].y + case 1: + prog = region.quantProgVals[hdr.quality].cb + default: + prog = region.quantProgVals[hdr.quality].cr + } + } + combined := progAdd(plain, prog) + var prev progBandQuant + switch which { + case 0: + prev = ts.yBitPos + case 1: + prev = ts.cbBitPos + default: + prev = ts.crBitPos + } + shift, ok := progLSub(combined, 1) + if !ok { + dumpUpgradeQuantOnce("lsub-underflow", which, plain, prog, prev, combined) + return progBandQuant{}, progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + numBits, ok := progSub(prev, combined) + if !ok { + dumpUpgradeQuantOnce("numbits-underflow", which, plain, prog, prev, combined) + return progBandQuant{}, progBandQuant{}, progBandQuant{}, progBandQuant{}, false + } + return plain, prog, shift, numBits, true +} + +// dumpUpgradeQuantOnce 诊断:UPGRADE 量化解析失败时转储全部带量化值。 +// 仅首条生效,避免刷屏。 +var upgradeQuantDumpOnce sync.Once + +func dumpUpgradeQuantOnce(kind string, which int, plain, prog, prev, combined progBandQuant) { + upgradeQuantDumpOnce.Do(func() { + slog.Warn("progressive: upgrade quant detail", + "kind", kind, "comp", which, + "plain", bandQuantStr(plain), + "prog", bandQuantStr(prog), + "prev", bandQuantStr(prev), + "combined", bandQuantStr(combined)) + }) +} + +func bandQuantStr(q progBandQuant) string { + return fmt.Sprintf("LL3=%d HL3=%d LH3=%d HH3=%d HL2=%d LH2=%d HH2=%d HL1=%d LH1=%d HH1=%d", + q.LL3, q.HL3, q.LH3, q.HH3, q.HL2, q.LH2, q.HH2, q.HL1, q.LH1, q.HH1) +} + +// progDecodeComponent implements progressive_rfx_decode_component for the +// SIMPLE/FIRST first pass: RLGR decode → sign snapshot → LL3 differential → +// per-band left-shift dequant → current update → IDWT. +// +// MS-RDPEGFX 3.3.8.2.1.1:LL3 差分累积与按带反量化(DecProgQ*PQF)对 +// ORIGINAL 与 DIFFERENCE 瓦片一视同仁;coeffDiff 只改变与参考状态 current +// 的合并方式(original 覆盖,difference 叠加)。与 FreeRDP +// progressive_rfx_decode_component / progressive_rfx_dwt_2d_decode 一致。 +func progDecodeComponent(data []byte, shift progBandQuant, buf, sign, current *coeffArr, coeffDiff, extrapolate bool) { + b := buf[:] + if data == nil || len(data) == 0 { + clear(b) + } else { + rlgr1Decode(data, 4096, b) + } + copy(sign[:], b) + + if !extrapolate { + progDiffDecode(b[4032:4096]) + progDecodeBlock(b[0:1024], shift.HL1) + progDecodeBlock(b[1024:2048], shift.LH1) + progDecodeBlock(b[2048:3072], shift.HH1) + progDecodeBlock(b[3072:3328], shift.HL2) + progDecodeBlock(b[3328:3584], shift.LH2) + progDecodeBlock(b[3584:3840], shift.HH2) + progDecodeBlock(b[3840:3904], shift.HL3) + progDecodeBlock(b[3904:3968], shift.LH3) + progDecodeBlock(b[3968:4032], shift.HH3) + progDecodeBlock(b[4032:4096], shift.LL3) + } else { + progDiffDecode(b[4015:4096]) + progDecodeBlock(b[0:1023], shift.HL1) + progDecodeBlock(b[1023:2046], shift.LH1) + progDecodeBlock(b[2046:3007], shift.HH1) + progDecodeBlock(b[3007:3279], shift.HL2) + progDecodeBlock(b[3279:3551], shift.LH2) + progDecodeBlock(b[3551:3807], shift.HH2) + progDecodeBlock(b[3807:3879], shift.HL3) + progDecodeBlock(b[3879:3951], shift.LH3) + progDecodeBlock(b[3951:4015], shift.HH3) + progDecodeBlock(b[4015:4096], shift.LL3) + } + + if coeffDiff { + for i := range b { + current[i] += b[i] + } + copy(b, current[:]) + } else { + copy(current[:], b) + } + if !extrapolate { + rfxInverseDWT2D(b) + } else { + progDWTExtrapolate(b) + } +} + +// progDiffDecode is rfx_differential_decode (in-place cumulative sum). +func progDiffDecode(data []int16) { + for i := 1; i < len(data); i++ { + data[i] += data[i-1] + } +} + +// progDecodeBlock is progressive_rfx_decode_block (left-shift dequant). +func progDecodeBlock(data []int16, shift uint8) { + if shift == 0 { + return + } + s := int16(shift) + for i := range data { + data[i] <<= s + } +} + +// ── Upgrade bit streams ──────────────────────────────────────────────────── + +// progBitStream is an MSB-first bit reader matching FreeRDP's wBitStream +// semantics (zero-padding past the end). +type progBitStream struct { + data []byte + bytePos int + acc uint32 + bits int + posBits int +} + +func (b *progBitStream) fill() { + for b.bits <= 24 && b.bytePos < len(b.data) { + b.acc |= uint32(b.data[b.bytePos]) << uint(24-b.bits) + b.bits += 8 + b.bytePos++ + } +} + +func (b *progBitStream) readBit() uint32 { + b.fill() + v := (b.acc >> 31) & 1 + b.acc <<= 1 + if b.bits > 0 { + b.bits-- + } + b.posBits++ + return v +} + +func (b *progBitStream) readBits(n uint) uint32 { + if n == 0 { + return 0 + } + b.fill() + var v uint32 + if b.bits >= int(n) { + v = (b.acc >> uint(32-int(n))) & uint32((1< 0 { + v = (b.acc >> uint(32-avail)) & uint32((1< 80 { + st.kp = 80 + } + st.nz-- + return 0 + } + // '1' bit: nz comes from the next k bits + st.nz = 0 + st.mode = 1 + if k > 0 { + st.nz = int(st.srl.readBits(k)) + } + if st.nz != 0 { + st.nz-- + return 0 + } + } + st.mode = 0 + // unary encoding; read sign bit + sign := st.srl.readBit() + if st.kp < 6 { + st.kp = 0 + } else { + st.kp -= 6 + } + if numBits == 1 { + if sign != 0 { + return -1 + } + return 1 + } + mag := uint32(1) + max := uint32(1< 32767 { + mag = 32767 + } + if sign != 0 { + return -int16(mag) + } + return int16(mag) +} + +func progRawShift(raw *progBitStream, numBits uint32) int16 { + return int16(raw.readBits(uint(numBits))) +} + +// progUpgradeBlock ports progressive_rfx_upgrade_block. +func progUpgradeBlock(st *progUpgradeState, buf, sign []int16, length uint32, shift, numBits uint32) { + if numBits < 1 { + return + } + raw := st.raw + if !st.nonLL { + for i := uint32(0); i < length; i++ { + input := progRawShift(raw, numBits) + buf[i] = int16(int32(buf[i]) + (int32(input) << shift)) + } + return + } + for i := uint32(0); i < length; i++ { + var input int32 + switch { + case sign[i] > 0: + input = int32(progRawShift(raw, numBits)) + case sign[i] < 0: + input = -int32(progRawShift(raw, numBits)) + default: + input = int32(st.srlRead(numBits)) + sign[i] = int16(input) + } + buf[i] = int16(int32(buf[i]) + (input << shift)) + } +} + +// progUpgradeStateFinish ports progressive_rfx_upgrade_state_finish: byte- +// align both streams and drop a trailing 8-bit srl remainder. +func progUpgradeStateFinish(st *progUpgradeState) { + raw, srl := st.raw, st.srl + if pad := (8 - raw.posBits%8) % 8; pad > 0 { + raw.skip(uint(pad)) + } + if pad := (8 - srl.posBits%8) % 8; pad > 0 { + srl.skip(uint(pad)) + } + if srl.remaining() == 8 { + srl.skip(8) + } +} + +// progUpgradeComponent ports progressive_rfx_upgrade_component: refines the +// cached coefficients (current) in extrapolate layout using an SRL stream +// (for sign==0 coefficients) and a RAW stream (for the rest). +func progUpgradeComponent(buf, current, sign *coeffArr, shift, numBits progBandQuant, srlData, rawData []byte, extrapolate bool) { + st := progUpgradeState{ + kp: 8, + mode: 0, + srl: &progBitStream{data: srlData}, + raw: &progBitStream{data: rawData}, + } + cur := current[:] + sgn := sign[:] + + st.nonLL = true + progUpgradeBlock(&st, cur[0:1023], sgn[0:1023], 1023, uint32(shift.HL1), uint32(numBits.HL1)) + progUpgradeBlock(&st, cur[1023:2046], sgn[1023:2046], 1023, uint32(shift.LH1), uint32(numBits.LH1)) + progUpgradeBlock(&st, cur[2046:3007], sgn[2046:3007], 961, uint32(shift.HH1), uint32(numBits.HH1)) + progUpgradeBlock(&st, cur[3007:3279], sgn[3007:3279], 272, uint32(shift.HL2), uint32(numBits.HL2)) + progUpgradeBlock(&st, cur[3279:3551], sgn[3279:3551], 272, uint32(shift.LH2), uint32(numBits.LH2)) + progUpgradeBlock(&st, cur[3551:3807], sgn[3551:3807], 256, uint32(shift.HH2), uint32(numBits.HH2)) + progUpgradeBlock(&st, cur[3807:3879], sgn[3807:3879], 72, uint32(shift.HL3), uint32(numBits.HL3)) + progUpgradeBlock(&st, cur[3879:3951], sgn[3879:3951], 72, uint32(shift.LH3), uint32(numBits.LH3)) + progUpgradeBlock(&st, cur[3951:4015], sgn[3951:4015], 64, uint32(shift.HH3), uint32(numBits.HH3)) + + st.nonLL = false + progUpgradeBlock(&st, cur[4015:4096], sgn[4015:4096], 81, uint32(shift.LL3), uint32(numBits.LL3)) + + progUpgradeStateFinish(&st) + + // dwt_2d_decode(..., reverse=TRUE): buffer = current, then IDWT. + copy(buf[:], cur) + if !extrapolate { + rfxInverseDWT2D(buf[:]) + } else { + progDWTExtrapolate(buf[:]) + } +} + +// ── Extrapolate IDWT (progressive_rfx_dwt_2d_decode_block) ───────────────── + +func progBandLCount(level int) int { return (64 >> level) + 1 } + +func progBandHCount(level int) int { + if level == 1 { + return (64 >> 1) - 1 + } + return (64 + (1 << uint(level-1))) >> level +} + +func progClamp16(v int32) int16 { + if v < -32768 { + return -32768 + } + if v > 32767 { + return 32767 + } + return int16(v) +} + +// progDWTExtrapolate ports rfx_dwt_2d_extrapolate_decode: three irregular +// blocks at fixed offsets covering the extrapolate band layout. +func progDWTExtrapolate(buffer []int16) { + bufs := idwtBufPool.Get().(*idwtBufs) + tmp := bufs.tmp[:] + progDWT2DBlock(buffer[3807:], tmp, 3) + progDWT2DBlock(buffer[3007:], tmp, 2) + progDWT2DBlock(buffer[0:], tmp, 1) + idwtBufPool.Put(bufs) +} + +// progDWT2DBlock decodes one extrapolate block in place. +func progDWT2DBlock(buffer, temp []int16, level int) { + nBandL := progBandLCount(level) + nBandH := progBandHCount(level) + + hlLen := nBandH * nBandL + lhLen := nBandL * nBandH + hhLen := nBandH * nBandH + llLen := nBandL * nBandL + + hl := buffer[0:hlLen] + lh := buffer[hlLen : hlLen+lhLen] + hh := buffer[hlLen+lhLen : hlLen+lhLen+hhLen] + ll := buffer[hlLen+lhLen+hhLen : hlLen+lhLen+hhLen+llLen] + + dstStep := nBandL + nBandH + lBuf := temp[0 : nBandL*dstStep] + hBuf := temp[nBandL*dstStep : nBandL*dstStep+nBandH*dstStep] + + progIDWTX(ll, nBandL, hl, nBandH, lBuf, dstStep, nBandL, nBandH, nBandL) + progIDWTX(lh, nBandL, hh, nBandH, hBuf, dstStep, nBandL, nBandH, nBandH) + progIDWTY(lBuf, dstStep, hBuf, dstStep, buffer, dstStep, nBandL, nBandH, nBandL+nBandH) +} + +// progIDWTX ports progressive_rfx_idwt_x (horizontal 1-D IDWT of every row). +// Index arithmetic instead of slice reslicing: the C original walks pointers +// one element past the final read, which Go bounds checks reject. +func progIDWTX(low []int16, lowStep int, high []int16, highStep int, dst []int16, dstStep int, lowCount, highCount, dstCount int) { + for i := 0; i < dstCount; i++ { + lRow := low[i*lowStep:] + hRow := high[i*highStep:] + xRow := dst[i*dstStep:] + H0 := hRow[0] + L0 := lRow[0] + li, hi := 1, 1 + xi := 0 + X0 := progClamp16(int32(L0) - int32(H0)) + X2 := X0 + for j := 0; j < highCount-1; j++ { + H1 := hRow[hi] + hi++ + L0 = lRow[li] + li++ + X2 = progClamp16(int32(L0) - (int32(H0)+int32(H1))/2) + X1 := progClamp16((int32(X0)+int32(X2))/2 + 2*int32(H0)) + xRow[xi] = X0 + xRow[xi+1] = X1 + xi += 2 + X0 = X2 + H0 = H1 + } + switch { + case lowCount <= highCount: + xRow[xi] = X2 + xRow[xi+1] = progClamp16(int32(X2) + 2*int32(H0)) + case lowCount == highCount+1: + L0 = lRow[li] + X0t := progClamp16(int32(L0) - int32(H0)) + xRow[xi] = X2 + xRow[xi+1] = progClamp16((int32(X0t)+int32(X2))/2 + 2*int32(H0)) + xRow[xi+2] = X0t + default: + L0 = lRow[li] + li++ + X0t := progClamp16(int32(L0) - int32(H0)/2) + xRow[xi] = X2 + xRow[xi+1] = progClamp16((int32(X0t)+int32(X2))/2 + 2*int32(H0)) + xRow[xi+2] = X0t + L0 = lRow[li] + xRow[xi+3] = progClamp16((int32(X0t) + int32(L0)) / 2) + } + } +} + +// progIDWTY ports progressive_rfx_idwt_y (vertical 1-D IDWT of every column). +// Index arithmetic instead of slice reslicing: the C original walks pointers +// one element past the final read, which Go bounds checks reject. +func progIDWTY(low []int16, lowStep int, high []int16, highStep int, dst []int16, dstStep int, lowCount, highCount, dstCount int) { + for i := 0; i < dstCount; i++ { + H0 := high[i] + L0 := low[i] + li, hi := 1, 1 + xi := 0 + X0 := progClamp16(int32(L0) - int32(H0)) + X2 := X0 + for j := 0; j < highCount-1; j++ { + H1 := high[i+hi*highStep] + hi++ + L0 = low[i+li*lowStep] + li++ + X2 = progClamp16(int32(L0) - (int32(H0)+int32(H1))/2) + X1 := progClamp16((int32(X0)+int32(X2))/2 + 2*int32(H0)) + dst[i+xi] = X0 + xi += dstStep + dst[i+xi] = X1 + xi += dstStep + X0 = X2 + H0 = H1 + } + switch { + case lowCount <= highCount: + dst[i+xi] = X2 + dst[i+xi+dstStep] = progClamp16(int32(X2) + 2*int32(H0)) + case lowCount == highCount+1: + L0 = low[i+li*lowStep] + X0t := progClamp16(int32(L0) - int32(H0)) + dst[i+xi] = X2 + dst[i+xi+dstStep] = progClamp16((int32(X0t)+int32(X2))/2 + 2*int32(H0)) + dst[i+xi+2*dstStep] = X0t + default: + L0 = low[i+li*lowStep] + li++ + X0t := progClamp16(int32(L0) - int32(H0)/2) + dst[i+xi] = X2 + dst[i+xi+dstStep] = progClamp16((int32(X0t)+int32(X2))/2 + 2*int32(H0)) + dst[i+xi+2*dstStep] = X0t + L0 = low[i+li*lowStep] + dst[i+xi+3*dstStep] = progClamp16((int32(X0t) + int32(L0)) / 2) + } + } +} + +// rfxPlaceTile converts YCbCr tile to BGRA using tile-grid indices (xIdx, yIdx). +// rfxPlaceTile 把解码后的瓦片绘制到表面。region.rects 非空时只绘制与矩形 +// 并集相交的瓦片,但相交的瓦片必须整块 64×64 落屏(FreeRDP update_tiles +// 语义:rects 仅用于筛掉不相交的瓦片,tile 边界按 64 对齐外延,矩形外的 +// 瓦片像素同样是本帧的有效内容)。实测 Win10 最小化动画:region 矩形 +// (34,254 662x374) 的瓦片从 yIdx=3(y=192)开始,让出的条带 y=198..254 +// 只存在于瓦片内——若按矩形交集裁剪,该条带永远不被重绘,留下残影。 +func rfxPlaceTile(yCoeffs, cbCoeffs, crCoeffs []int16, xIdx, yIdx int, output []byte, outW, outH int, rects []rfxRect) { + tileX := xIdx * rfxTileSize + tileY := yIdx * rfxTileSize + if len(rects) > 0 { + intersects := false + for _, rc := range rects { + if tileX < rc.x+rc.w && tileX+rfxTileSize > rc.x && + tileY < rc.y+rc.h && tileY+rfxTileSize > rc.y { + intersects = true + break + } + } + if !intersects { + return + } + } + rfxPlaceTileAbs(yCoeffs, cbCoeffs, crCoeffs, tileX, tileY, output, outW, outH) +} + +// rfxPlaceTileAbs converts YCbCr tile to BGRA and writes into the output buffer +// at absolute pixel coordinates (tileX, tileY). +// Uses ICT (Irreversible Color Transform) from MS-RDPRFX. +func rfxPlaceTileAbs(yCoeffs, cbCoeffs, crCoeffs []int16, tileX, tileY int, output []byte, outW, outH int) { + tileW := rfxTileSize + tileH := rfxTileSize + if tileX+tileW > outW { + tileW = outW - tileX + } + if tileY+tileH > outH { + tileH = outH - tileY + } + if tileW <= 0 || tileH <= 0 { + return + } + + for row := 0; row < tileH; row++ { + dstStart := ((tileY+row)*outW + tileX) * 4 + dstEnd := dstStart + tileW*4 + if dstStart < 0 || dstEnd > len(output) { + continue + } + dstRow := output[dstStart:dstEnd:dstEnd] + srcOff := row * rfxTileSize + ictToBGRA( + yCoeffs[srcOff:srcOff+tileW:srcOff+tileW], + cbCoeffs[srcOff:srcOff+tileW:srcOff+tileW], + crCoeffs[srcOff:srcOff+tileW:srcOff+tileW], + dstRow, tileW, + ) + } +} diff --git a/plugin/rdpgfx/rfx_progressive_test.go b/plugin/rdpgfx/rfx_progressive_test.go new file mode 100644 index 0000000..c1a6817 --- /dev/null +++ b/plugin/rdpgfx/rfx_progressive_test.go @@ -0,0 +1,57 @@ +package rdpgfx + +import ( + "testing" +) + +// BenchmarkIctToBGRA benchmarks ictToBGRA on a full 64×64 tile row (64 pixels). +func BenchmarkIctToBGRA(b *testing.B) { + n := rfxTileSize // 64 pixels per row + yRow := make([]int16, n) + cbRow := make([]int16, n) + crRow := make([]int16, n) + dst := make([]byte, n*4) + for i := range n { + yRow[i] = int16(i * 4) + cbRow[i] = int16(i%64 - 32) + crRow[i] = int16(i%32 - 16) + } + b.ResetTimer() + for b.Loop() { + ictToBGRA(yRow, cbRow, crRow, dst, n) + } +} + +// BenchmarkRfxDecodeComponent benchmarks a full component decode pipeline: +// RLGR → differential/dequantize → inverse DWT. +func BenchmarkRfxDecodeComponent(b *testing.B) { + // Use the same coefficient pattern as the RLGR benchmarks (1/17 non-zero). + coeffs := make([]int16, 4096) + for i := range coeffs { + if i%17 == 0 { + coeffs[i] = int16(i%256 - 128) + } + } + data := rlgr1Encode(coeffs) + quant := rfxQuant{6, 6, 6, 6, 6, 6, 6, 6, 6, 6} + + b.ResetTimer() + var dst []int16 + for b.Loop() { + dst = rfxDecodeComponent(data, quant, 1) + coeffPool.Put((*coeffArr)(dst)) + dst = nil + } +} + +// BenchmarkRfxInverseDWT2D benchmarks the full 3-level inverse DWT on 4096 coefficients. +func BenchmarkRfxInverseDWT2D(b *testing.B) { + coeffs := make([]int16, 4096) + for i := range coeffs { + coeffs[i] = int16(i%64 - 32) + } + b.ResetTimer() + for b.Loop() { + rfxInverseDWT2D(coeffs) + } +} diff --git a/plugin/rdpgfx/rfx_rlgr.go b/plugin/rdpgfx/rfx_rlgr.go new file mode 100644 index 0000000..815cb25 --- /dev/null +++ b/plugin/rdpgfx/rfx_rlgr.go @@ -0,0 +1,452 @@ +package rdpgfx + +// RLGR1/RLGR3 (Run-Length Golomb-Rice) decoder for RFX codec. +// Reference: MS-RDPRFX 3.1.8.1.7.3 RLGR1/RLGR3 Pseudocode +// Matches FreeRDP's rfx_rlgr.c implementation. + +import "math/bits" + +const ( + rlgrLSGR = 3 // shift count to convert kp to k + rlgrKPMax = 80 // max value for kp or krp + rlgrUPGR = 4 // increase in kp after a zero run in RL mode + rlgrDNGR = 6 // decrease in kp after a nonzero symbol in RL mode + rlgrUQGR = 3 // increase in kp after zero symbol in GR mode + rlgrDQGR = 3 // decrease in kp after nonzero symbol in GR mode +) + +// rlgr1Decode decodes RLGR1-encoded data into signed 16-bit DWT coefficients. +// If dst is non-nil and has sufficient capacity, it is reused (zeroed first). +func rlgr1Decode(data []byte, outputSize int, dst []int16) []int16 { + var output []int16 + if cap(dst) >= outputSize { + output = dst[:outputSize] + clear(output) + } else { + output = make([]int16, outputSize) + } + br := &rlgrBitReader{data: data} + cnt := 0 + + k := uint32(1) + kp := uint32(1 << rlgrLSGR) // 8 + kr := uint32(1) + krp := uint32(1 << rlgrLSGR) // 8 + + for br.remaining() > 0 && cnt < outputSize { + if k > 0 { + // RL (Run-Length) Mode + + // Count leading 0-bits → number of full run groups + vk := br.countLeadingZeros() + + // Each leading 0 adds (1 << k) to run, with k adapting upward + run := uint32(0) + for range vk { + run += 1 << k + kp += rlgrUPGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + } + + // Read k bits for run remainder + if k > 0 { + run += br.readBits(int(k)) + } + + // Read sign bit for the non-zero value + sign := br.readBits(1) + + // Decode non-zero magnitude using GR code with leading 1-bits + vk2 := br.countLeadingOnes() + + // Read kr bits for code remainder + code := uint32(0) + if kr > 0 { + code = br.readBits(int(kr)) + } + code |= vk2 << kr + + // Update kr/krp + if vk2 == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk2 != 1 { + krp += vk2 + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + // Update k/kp (decrease after non-zero) + if kp > rlgrDNGR { + kp -= rlgrDNGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + + // Compute magnitude (code + 1, guaranteed non-zero) + mag := int16(code + 1) + if sign != 0 { + mag = -mag + } + + // Output: run zeros (already 0 from init), then the non-zero value + runEnd := min(cnt+int(run), outputSize) + cnt = runEnd + if cnt < outputSize { + output[cnt] = mag + cnt++ + } + + } else { + // GR (Golomb-Rice) Mode + + // Count leading 1-bits + vk := br.countLeadingOnes() + + // Read kr bits for code remainder + code := uint32(0) + if kr > 0 { + code = br.readBits(int(kr)) + } + code |= vk << kr + + // Update kr/krp + if vk == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk != 1 { + krp += vk + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + // RLGR1: sign embedded in code as code = 2*magnitude - sign + if code == 0 { + kp += rlgrUQGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + + if cnt < outputSize { + cnt++ // zero already set from init + } + } else { + if kp > rlgrDQGR { + kp -= rlgrDQGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + + var mag int16 + if code&1 != 0 { + // odd code → negative + mag = -int16((code + 1) >> 1) + } else { + // even code → positive + mag = int16(code >> 1) + } + if cnt < outputSize { + output[cnt] = mag + cnt++ + } + } + } + } + + return output +} + +// rlgr3Decode decodes RLGR3-encoded data into signed 16-bit DWT coefficients. +// RLGR3 differs from RLGR1 only in GR mode: it encodes/decodes TWO values +// per GR code by encoding their sum then splitting. +// Reference: MS-RDPRFX 3.1.8.1.7.3, FreeRDP rfx_rlgr.c +func rlgr3Decode(data []byte, outputSize int, dst []int16) []int16 { + var output []int16 + if cap(dst) >= outputSize { + output = dst[:outputSize] + clear(output) + } else { + output = make([]int16, outputSize) + } + br := &rlgrBitReader{data: data} + cnt := 0 + + k := uint32(1) + kp := uint32(1 << rlgrLSGR) + kr := uint32(1) + krp := uint32(1 << rlgrLSGR) + + for br.remaining() > 0 && cnt < outputSize { + if k > 0 { + // RL Mode — identical to RLGR1 + vk := br.countLeadingZeros() + + run := uint32(0) + for range vk { + run += 1 << k + kp += rlgrUPGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + } + + if k > 0 { + run += br.readBits(int(k)) + } + + sign := br.readBits(1) + + vk2 := br.countLeadingOnes() + + code := uint32(0) + if kr > 0 { + code = br.readBits(int(kr)) + } + code |= vk2 << kr + + if vk2 == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk2 != 1 { + krp += vk2 + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + if kp > rlgrDNGR { + kp -= rlgrDNGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + + mag := int16(code + 1) + if sign != 0 { + mag = -mag + } + + runEnd3 := min(cnt+int(run), outputSize) + cnt = runEnd3 + if cnt < outputSize { + output[cnt] = mag + cnt++ + } + + } else { + // GR Mode — RLGR3 variant: decode TWO values from one GR code + vk := br.countLeadingOnes() + + code := uint32(0) + if kr > 0 { + code = br.readBits(int(kr)) + } + code |= vk << kr + + if vk == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk != 1 { + krp += vk + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + // RLGR3: code = val1 + val2 (sum of two 2*mag-sign encoded values) + // Read nIdx bits to split: nIdx = bit-length of code + nIdx := uint32(0) + if code != 0 { + nIdx = uint32(bits.Len(uint(code))) + } + + if br.remaining() < int(nIdx) { + break + } + val1 := uint32(0) + if nIdx > 0 { + val1 = br.readBits(int(nIdx)) + } + val2 := code - val1 + + // Update k/kp based on both values + if val1 != 0 && val2 != 0 { + if kp > 2*rlgrDQGR { + kp -= 2 * rlgrDQGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + } else if val1 == 0 && val2 == 0 { + kp += 2 * rlgrUQGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + } + + // Decode val1 as 2*mag-sign + var mag1 int16 + if val1&1 != 0 { + mag1 = -int16((val1 + 1) >> 1) + } else { + mag1 = int16(val1 >> 1) + } + if cnt < outputSize { + output[cnt] = mag1 + cnt++ + } + + // Decode val2 as 2*mag-sign + var mag2 int16 + if val2&1 != 0 { + mag2 = -int16((val2 + 1) >> 1) + } else { + mag2 = int16(val2 >> 1) + } + if cnt < outputSize { + output[cnt] = mag2 + cnt++ + } + } + } + + return output +} + +// rlgrBitReader reads bits MSB-first from a byte slice. +// +// To minimise the per-bit cost on the RLGR hot path we keep a 64-bit +// shift-register (`acc`, MSB-aligned with `bitsInAcc` valid bits at the top) +// fed from `data[bytePos:]`. Reads up to 32 bits are a shift+mask, and runs +// of identical bits are extracted with a single `bits.LeadingZeros64`. +// +// Invariant: bits consumed == bytePos*8 - bitsInAcc, so +// +// remaining() == (len(data)-bytePos)*8 + bitsInAcc +// +// which lets us drop the separate `total` and `read` counters entirely. +type rlgrBitReader struct { + data []byte + bytePos int // next byte to load into acc + acc uint64 // bits aligned to MSB + bitsInAcc int // number of valid bits in acc (MSB-aligned) +} + +func (br *rlgrBitReader) remaining() int { + return (len(br.data)-br.bytePos)*8 + br.bitsInAcc +} + +// fill loads bytes into the high end of acc until at least `need` bits are +// buffered or the input is exhausted. need must be <= 56. +func (br *rlgrBitReader) fill(need int) { + for br.bitsInAcc < need && br.bytePos < len(br.data) { + br.acc |= uint64(br.data[br.bytePos]) << uint(56-br.bitsInAcc) + br.bytePos++ + br.bitsInAcc += 8 + } +} + +// readBits extracts n bits (n > 0) from the accumulator. +// Inlinable: when bitsInAcc is already sufficient the slow path is never +// compiled into the call site; when fill is also inlinable the whole hot +// path reduces to a shift + mask without a call frame. +func (br *rlgrBitReader) readBits(n int) uint32 { + if br.bitsInAcc < n { + br.fill(n) + if br.bitsInAcc < n { + // EOF (bytePos == len(data)): zero bitsInAcc so remaining() + // returns 0 and the decode loop terminates cleanly. + br.bitsInAcc = 0 + return 0 + } + } + val := uint32(br.acc >> uint(64-n)) + br.acc <<= uint(n) + br.bitsInAcc -= n + return val +} + +// countLeadingZeros counts consecutive 0-bits and consumes the first 1-bit terminator. +func (br *rlgrBitReader) countLeadingZeros() uint32 { + count := uint32(0) + for { + if br.bitsInAcc < 56 && br.bytePos < len(br.data) { + br.fill(56) + } + if br.bitsInAcc == 0 { + return count + } + lz := bits.LeadingZeros64(br.acc) + if lz >= br.bitsInAcc { + count += uint32(br.bitsInAcc) + br.acc = 0 + br.bitsInAcc = 0 + continue + } + count += uint32(lz) + consume := lz + 1 + br.acc <<= uint(consume) + br.bitsInAcc -= consume + return count + } +} + +// countLeadingOnes counts consecutive 1-bits and consumes the first 0-bit terminator. +func (br *rlgrBitReader) countLeadingOnes() uint32 { + count := uint32(0) + for { + if br.bitsInAcc < 56 && br.bytePos < len(br.data) { + br.fill(56) + } + if br.bitsInAcc == 0 { + return count + } + lo := bits.LeadingZeros64(^br.acc) + if lo >= br.bitsInAcc { + count += uint32(br.bitsInAcc) + br.acc = 0 + br.bitsInAcc = 0 + continue + } + count += uint32(lo) + consume := lo + 1 + br.acc <<= uint(consume) + br.bitsInAcc -= consume + return count + } +} + +// DecodeRLGR3ForDebug exposes the RLGR3 decoder for offline verification. +func DecodeRLGR3ForDebug(data []byte, outputSize int) []int16 { + return rlgr3Decode(data, outputSize, nil) +} diff --git a/plugin/rdpgfx/rfx_rlgr_test.go b/plugin/rdpgfx/rfx_rlgr_test.go new file mode 100644 index 0000000..a71b199 --- /dev/null +++ b/plugin/rdpgfx/rfx_rlgr_test.go @@ -0,0 +1,581 @@ +package rdpgfx + +import ( + "testing" +) + +// rlgr1Encode is a minimal RLGR1 encoder used only for round-trip tests. +// It implements the exact inverse of rlgr1Decode per MS-RDPRFX 3.1.8.1.7.3. +func rlgr1Encode(coeffs []int16) []byte { + bw := &bitWriter{} + k := uint32(1) + kp := uint32(1 << rlgrLSGR) + kr := uint32(1) + krp := uint32(1 << rlgrLSGR) + + i := 0 + for i < len(coeffs) { + if k > 0 { + // RL mode: count zeros then encode non-zero value + numZeros := uint32(0) + for i+int(numZeros) < len(coeffs) && coeffs[i+int(numZeros)] == 0 { + numZeros++ + } + + runMax := uint32(1) << k + nGroups := uint32(0) + run := numZeros + for run >= runMax { + bw.writeBit(0) // leading 0 + run -= runMax + kp += rlgrUPGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + runMax = 1 << k + nGroups++ + } + bw.writeBit(1) // terminator + if k > 0 { + bw.writeBits(run, int(k)) + } + i += int(numZeros) + + if i >= len(coeffs) { + break + } + + // Encode non-zero value + val := coeffs[i] + i++ + + sign := uint32(0) + mag := uint32(val) + if val < 0 { + sign = 1 + mag = uint32(-val) + } + bw.writeBit(sign) // sign bit + + code := mag - 1 + vk2 := code >> kr + remainder := code & ((1 << kr) - 1) + + // leading 1-bits + for range vk2 { + bw.writeBit(1) + } + bw.writeBit(0) // terminator + if kr > 0 { + bw.writeBits(remainder, int(kr)) + } + + // Update kr/krp + if vk2 == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk2 != 1 { + krp += vk2 + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + // Update k/kp + if kp > rlgrDNGR { + kp -= rlgrDNGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + + } else { + // GR mode: encode single value + val := coeffs[i] + i++ + + var code uint32 + if val == 0 { + code = 0 + kp += rlgrUQGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + } else { + if val > 0 { + code = uint32(val) * 2 + } else { + code = uint32(-val)*2 - 1 + } + if kp > rlgrDQGR { + kp -= rlgrDQGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + } + + vk := code >> kr + remainder := code & ((1 << kr) - 1) + + // leading 1-bits + for range vk { + bw.writeBit(1) + } + bw.writeBit(0) // terminator + if kr > 0 { + bw.writeBits(remainder, int(kr)) + } + + // Update kr/krp + if vk == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk != 1 { + krp += vk + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + } + } + + return bw.bytes() +} + +type bitWriter struct { + data []byte + bitPos int // 0..7, bits written in current byte + current byte +} + +func (bw *bitWriter) writeBit(b uint32) { + bw.current = (bw.current << 1) | byte(b&1) + bw.bitPos++ + if bw.bitPos == 8 { + bw.data = append(bw.data, bw.current) + bw.current = 0 + bw.bitPos = 0 + } +} + +func (bw *bitWriter) writeBits(val uint32, n int) { + for i := n - 1; i >= 0; i-- { + bw.writeBit((val >> uint(i)) & 1) + } +} + +func (bw *bitWriter) bytes() []byte { + if bw.bitPos > 0 { + bw.data = append(bw.data, bw.current< 10 { + t.Fatalf("Too many errors, stopping") + } + } + } + t.Logf("LL3[0]=%d (expected %d)", decoded[4032], tc.dc) + }) + } +} + +func TestRLGR1RoundTrip_MultipleNonZero(t *testing.T) { + // Coefficients with non-zero values in various subbands + coeffs := make([]int16, 4096) + coeffs[0] = 5 // HL1[0] + coeffs[100] = -3 // HL1[100] + coeffs[1024] = 7 // LH1[0] + coeffs[4032] = 10 // LL3[0] + coeffs[4033] = 2 // LL3[1] + + encoded := rlgr1Encode(coeffs) + t.Logf("Multi non-zero: encoded to %d bytes", len(encoded)) + + decoded := rlgr1Decode(encoded, 4096, nil) + + for i, v := range decoded { + if v != coeffs[i] { + t.Errorf("Position %d: expected %d, got %d", i, coeffs[i], v) + } + } +} + +func TestRLGR1Decode_KnownBytes(t *testing.T) { + // Test with the all-zero Cb/Cr data (5 bytes) that the server sends + // This should decode to all zeros + // We'll encode all zeros and verify the decoder handles it + coeffs := make([]int16, 4096) + encoded := rlgr1Encode(coeffs) + + decoded := rlgr1Decode(encoded, 4096, nil) + for i, v := range decoded { + if v != 0 { + t.Errorf("Position %d: expected 0, got %d", i, v) + } + } +} + +// rlgr3Encode is a minimal RLGR3 encoder used only for round-trip tests. +// RLGR3 differs from RLGR1 only in GR mode: it encodes TWO values per code. +// Reference: FreeRDP rfx_rlgr.c (rfx_rlgr3_encode). +func rlgr3Encode(coeffs []int16) []byte { + bw := &bitWriter{} + k := uint32(1) + kp := uint32(1 << rlgrLSGR) + kr := uint32(1) + krp := uint32(1 << rlgrLSGR) + + i := 0 + for i < len(coeffs) { + if k > 0 { + // RL mode: identical to RLGR1 + numZeros := uint32(0) + for i+int(numZeros) < len(coeffs) && coeffs[i+int(numZeros)] == 0 { + numZeros++ + } + + runMax := uint32(1) << k + nGroups := uint32(0) + run := numZeros + for run >= runMax { + bw.writeBit(0) + run -= runMax + nGroups++ + kp += rlgrUPGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + runMax = uint32(1) << k + } + bw.writeBit(1) // terminator + if k > 0 { + bw.writeBits(run, int(k)) + } + i += int(numZeros) + _ = nGroups + + if i >= len(coeffs) { + break + } + + // Encode the non-zero value with sign + val := coeffs[i] + i++ + + var code uint32 + sign := uint32(0) + if val < 0 { + code = uint32(-val) - 1 + sign = 1 + } else { + code = uint32(val) - 1 + } + bw.writeBit(sign) + + vk := code >> kr + remainder := code & ((1 << kr) - 1) + for range vk { + bw.writeBit(1) + } + bw.writeBit(0) + if kr > 0 { + bw.writeBits(remainder, int(kr)) + } + + if vk == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk != 1 { + krp += vk + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + kp -= rlgrDNGR + if kp > rlgrKPMax { // underflow + kp = 0 + } + k = kp >> rlgrLSGR + + } else { + // GR mode: RLGR3 encodes TWO values per code + var val1, val2 int16 + val1 = coeffs[i] + i++ + if i < len(coeffs) { + val2 = coeffs[i] + i++ + } + + // Convert to 2*mag-sign encoding + var u1, u2 uint32 + if val1 < 0 { + u1 = uint32(-val1)*2 - 1 + } else { + u1 = uint32(val1) * 2 + } + if val2 < 0 { + u2 = uint32(-val2)*2 - 1 + } else { + u2 = uint32(val2) * 2 + } + + code := u1 + u2 + + // GR encode the sum + vk := code >> kr + remainder := code & ((1 << kr) - 1) + for range vk { + bw.writeBit(1) + } + bw.writeBit(0) + if kr > 0 { + bw.writeBits(remainder, int(kr)) + } + + if vk == 0 { + if krp > 2 { + krp -= 2 + } else { + krp = 0 + } + kr = krp >> rlgrLSGR + } else if vk != 1 { + krp += vk + if krp > rlgrKPMax { + krp = rlgrKPMax + } + kr = krp >> rlgrLSGR + } + + // Write nIdx bits for val1 + nIdx := uint32(0) + if code != 0 { + nIdx = uint32(bitLen(code)) + } + if nIdx > 0 { + bw.writeBits(u1, int(nIdx)) + } + + // Update k/kp + if u1 != 0 && u2 != 0 { + if kp > 2*rlgrDQGR { + kp -= 2 * rlgrDQGR + } else { + kp = 0 + } + k = kp >> rlgrLSGR + } else if u1 == 0 && u2 == 0 { + kp += 2 * rlgrUQGR + if kp > rlgrKPMax { + kp = rlgrKPMax + } + k = kp >> rlgrLSGR + } + } + } + + return bw.bytes() +} + +// bitLen returns the number of bits needed to represent val (same as bits.Len). +func bitLen(val uint32) int { + n := 0 + for val > 0 { + n++ + val >>= 1 + } + return n +} + +func TestRLGR3RoundTrip_AllZeros(t *testing.T) { + coeffs := make([]int16, 4096) + encoded := rlgr3Encode(coeffs) + t.Logf("All zeros: encoded to %d bytes", len(encoded)) + + decoded := rlgr3Decode(encoded, 4096, nil) + for i, v := range decoded { + if v != 0 { + t.Errorf("Position %d: expected 0, got %d", i, v) + } + } +} + +func TestRLGR3RoundTrip_SingleDC(t *testing.T) { + testCases := []struct { + name string + dc int16 + }{ + {"DC=+3", 3}, + {"DC=-4", -4}, + {"DC=+127", 127}, + {"DC=-128", -128}, + {"DC=+1", 1}, + {"DC=-1", -1}, + {"DC=+50", 50}, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + coeffs := make([]int16, 4096) + coeffs[4032] = tc.dc + + encoded := rlgr3Encode(coeffs) + t.Logf("Encoded to %d bytes (hex: % x)", len(encoded), encoded) + + decoded := rlgr3Decode(encoded, 4096, nil) + + for i, v := range decoded { + if v != coeffs[i] { + t.Errorf("Position %d: expected %d, got %d", i, coeffs[i], v) + if i > 10 { + t.Fatalf("Too many errors, stopping") + } + } + } + t.Logf("LL3[0]=%d (expected %d)", decoded[4032], tc.dc) + }) + } +} + +func TestRLGR3RoundTrip_MultipleNonZero(t *testing.T) { + coeffs := make([]int16, 4096) + coeffs[0] = 5 + coeffs[1] = -3 + coeffs[100] = 7 + coeffs[101] = -2 + coeffs[1024] = 10 + coeffs[4032] = 20 + coeffs[4033] = -15 + + encoded := rlgr3Encode(coeffs) + t.Logf("Multi non-zero: encoded to %d bytes", len(encoded)) + + decoded := rlgr3Decode(encoded, 4096, nil) + + for i, v := range decoded { + if v != coeffs[i] { + t.Errorf("Position %d: expected %d, got %d", i, coeffs[i], v) + } + } +} + +func TestRLGR3RoundTrip_ConsecutiveNonZero(t *testing.T) { + // Test GR mode with many consecutive non-zero pairs + coeffs := make([]int16, 64) + for i := range 64 { + coeffs[i] = int16(i%7 - 3) // values: -3,-2,-1,0,1,2,3 + } + + encoded := rlgr3Encode(coeffs) + t.Logf("Consecutive non-zero: encoded to %d bytes", len(encoded)) + + decoded := rlgr3Decode(encoded, 64, nil) + + for i, v := range decoded { + if v != coeffs[i] { + t.Errorf("Position %d: expected %d, got %d", i, coeffs[i], v) + } + } +} + +// BenchmarkRLGR1Decode benchmarks RLGR1 decode on a 64x64 coefficient block. +func BenchmarkRLGR1Decode(b *testing.B) { + // Create coefficients representative of a natural image tile (mostly zeros, some non-zero) + coeffs := make([]int16, 4096) + for i := range coeffs { + if i%17 == 0 { + coeffs[i] = int16(i%256 - 128) + } + } + data := rlgr1Encode(coeffs) + + b.ResetTimer() + var dst []int16 + for b.Loop() { + dst = rlgr1Decode(data, len(coeffs), dst) + } + _ = dst +} + +// BenchmarkRLGR3Decode benchmarks RLGR3 decode on a 64x64 coefficient block. +func BenchmarkRLGR3Decode(b *testing.B) { + coeffs := make([]int16, 4096) + for i := range coeffs { + if i%17 == 0 { + coeffs[i] = int16(i%256 - 128) + } + } + data := rlgr3Encode(coeffs) + + b.ResetTimer() + var dst []int16 + for b.Loop() { + dst = rlgr3Decode(data, len(coeffs), dst) + } + _ = dst +} diff --git a/plugin/rdpgfx/zgfx.go b/plugin/rdpgfx/zgfx.go new file mode 100644 index 0000000..8642cd3 --- /dev/null +++ b/plugin/rdpgfx/zgfx.go @@ -0,0 +1,493 @@ +package rdpgfx + +// ZGFX (RDP8 Bulk Compression) decompressor. +// Implements the decompression algorithm described in MS-RDPEGFX 2.2.4 / 3.3.8. +// Based on FreeRDP's reference implementation (libfreerdp/codec/zgfx.c). + +const zgfxHistorySize = 2500000 + +type zgfxContext struct { + history []byte + historyIdx int +} + +func newZgfxContext() *zgfxContext { + return &zgfxContext{ + history: make([]byte, zgfxHistorySize), + } +} + +// Huffman token types +const ( + tokenLiteral = 0 + tokenMatch = 1 +) + +type zgfxToken struct { + prefixLen uint8 + prefixCode uint16 + valueBits uint8 + tokenType uint8 + valueBase uint32 +} + +// Fixed Huffman table from FreeRDP zgfx.c (Apache-2.0 licensed). +var zgfxTokenTable = []zgfxToken{ + {1, 0, 8, tokenLiteral, 0}, // 0 + {5, 17, 5, tokenMatch, 0}, // 10001 + {5, 18, 7, tokenMatch, 32}, // 10010 + {5, 19, 9, tokenMatch, 160}, // 10011 + {5, 20, 10, tokenMatch, 672}, // 10100 + {5, 21, 12, tokenMatch, 1696}, // 10101 + {5, 24, 0, tokenLiteral, 0x00}, // 11000 + {5, 25, 0, tokenLiteral, 0x01}, // 11001 + {6, 44, 14, tokenMatch, 5792}, // 101100 + {6, 45, 15, tokenMatch, 22176}, // 101101 + {6, 52, 0, tokenLiteral, 0x02}, // 110100 + {6, 53, 0, tokenLiteral, 0x03}, // 110101 + {6, 54, 0, tokenLiteral, 0xFF}, // 110110 + {7, 92, 18, tokenMatch, 54944}, // 1011100 + {7, 93, 20, tokenMatch, 317088}, // 1011101 + {7, 110, 0, tokenLiteral, 0x04}, // 1101110 + {7, 111, 0, tokenLiteral, 0x05}, // 1101111 + {7, 112, 0, tokenLiteral, 0x06}, // 1110000 + {7, 113, 0, tokenLiteral, 0x07}, // 1110001 + {7, 114, 0, tokenLiteral, 0x08}, // 1110010 + {7, 115, 0, tokenLiteral, 0x09}, // 1110011 + {7, 116, 0, tokenLiteral, 0x0A}, // 1110100 + {7, 117, 0, tokenLiteral, 0x0B}, // 1110101 + {7, 118, 0, tokenLiteral, 0x3A}, // 1110110 + {7, 119, 0, tokenLiteral, 0x3B}, // 1110111 + {7, 120, 0, tokenLiteral, 0x3C}, // 1111000 + {7, 121, 0, tokenLiteral, 0x3D}, // 1111001 + {7, 122, 0, tokenLiteral, 0x3E}, // 1111010 + {7, 123, 0, tokenLiteral, 0x3F}, // 1111011 + {7, 124, 0, tokenLiteral, 0x40}, // 1111100 + {7, 125, 0, tokenLiteral, 0x80}, // 1111101 + {8, 188, 20, tokenMatch, 1365664}, // 10111100 + {8, 189, 21, tokenMatch, 2414240}, // 10111101 + {8, 252, 0, tokenLiteral, 0x0C}, // 11111100 + {8, 253, 0, tokenLiteral, 0x38}, // 11111101 + {8, 254, 0, tokenLiteral, 0x39}, // 11111110 + {8, 255, 0, tokenLiteral, 0x66}, // 11111111 + {9, 380, 22, tokenMatch, 4511392}, // 101111100 + {9, 381, 23, tokenMatch, 8705696}, // 101111101 + {9, 382, 24, tokenMatch, 17094304}, // 101111110 +} + +// zgfxTokenLut maps the next 9 bits (MSB-first) of the input stream to the +// matching Huffman token entry. Built once at init(). Replaces the per-bit +// linear scan over zgfxTokenTable that used to dominate ZGFX hot path +// profiles — typical RDPGFX traffic decodes thousands of tokens per frame. +type tokenLutEntry struct { + prefixLen uint8 // 0 means "no token (truncated input)" + valueBits uint8 + tokenType uint8 + valueBase uint32 +} + +var zgfxTokenLut [512]tokenLutEntry + +func init() { + for _, t := range zgfxTokenTable { + // All 9-bit windows whose top prefixLen bits == prefixCode map to t. + shift := uint(9 - t.prefixLen) + base := uint32(t.prefixCode) << shift + span := uint32(1) << shift + for j := range span { + idx := base | j + if zgfxTokenLut[idx].prefixLen == 0 { + zgfxTokenLut[idx] = tokenLutEntry{ + prefixLen: t.prefixLen, + valueBits: t.valueBits, + tokenType: t.tokenType, + valueBase: t.valueBase, + } + } + } + } +} + +// bitReader reads bits MSB-first from a byte slice. +type bitReader struct { + data []byte + bytePos int + bitPos uint8 // bits remaining in current byte (8..1) + bitsRemaining uint32 // total decodable bits remaining +} + +func newBitReader(data []byte) *bitReader { + br := &bitReader{data: data} + if len(data) > 0 { + br.bitPos = 8 + } + return br +} + +// newBitReaderWithCount creates a reader that tracks total decodable bits. +// The last byte of RDP8 compressed data encodes the number of padding bits +// to subtract: bitsAvailable = (len-1)*8 - lastByte. +func newBitReaderWithCount(data []byte) *bitReader { + br := &bitReader{data: data} + if len(data) < 2 { + return br + } + br.bitPos = 8 + paddingBits := uint32(data[len(data)-1]) + totalBits := uint32(len(data)-1) * 8 + if paddingBits > totalBits { + br.bitsRemaining = 0 + } else { + br.bitsRemaining = totalBits - paddingBits + } + // Exclude the last byte from readable data + br.data = data[:len(data)-1] + return br +} + +func (br *bitReader) hasBitsRemaining() bool { + return br.bitsRemaining > 0 +} + +// peek9 returns up to 9 bits left-aligned as the high bits of a 9-bit value +// (i.e. bit 8 = next bit out of the stream). Does not advance the reader. +// avail is the number of valid bits returned; if fewer than 9 bits are +// available the returned word is zero-padded on the right (LSB). +func (br *bitReader) peek9() (val uint32, avail uint8) { + if br.bytePos >= len(br.data) { + return 0, 0 + } + // First byte: only the low br.bitPos bits are still unread. + mask := uint32(1)<= 9 { + val = (bits >> (uint(avail) - 9)) & 0x1FF + avail = 9 + } else if avail > 0 { + val = bits << (9 - avail) + } + // Then clamp the reported avail to bitsRemaining. val keeps its high + // bits valid; lower bits become don't-care padding. decodeToken only + // accepts a LUT entry whose prefixLen <= the (clamped) avail, so any + // bits beyond bitsRemaining cannot influence the decoded prefix. + if uint32(avail) > br.bitsRemaining { + avail = uint8(br.bitsRemaining) + } + return +} + +// consumeBits advances the reader by n bits, updating bytePos/bitPos and +// bitsRemaining. Caller must ensure n <= bitsRemaining (LUT entries +// already include this guarantee for valid streams). +func (br *bitReader) consumeBits(n uint8) { + if uint32(n) > br.bitsRemaining { + n = uint8(br.bitsRemaining) + } + br.bitsRemaining -= uint32(n) + for n >= br.bitPos { + n -= br.bitPos + br.bytePos++ + br.bitPos = 8 + } + br.bitPos -= n + if br.bitPos == 0 { + br.bytePos++ + br.bitPos = 8 + } +} + +func (br *bitReader) getBit() uint32 { + if br.bytePos >= len(br.data) { + return 0 + } + br.bitPos-- + bit := uint32((br.data[br.bytePos] >> br.bitPos) & 1) + if br.bitPos == 0 { + br.bytePos++ + br.bitPos = 8 + } + return bit +} + +func (br *bitReader) getBits(n uint8) uint32 { + if n == 0 { + return 0 + } + // Fast path: enough bits remain in the current byte. + if n <= br.bitPos && br.bytePos < len(br.data) { + br.bitPos -= n + v := (uint32(br.data[br.bytePos]) >> br.bitPos) & ((1 << n) - 1) + if br.bitPos == 0 { + br.bytePos++ + br.bitPos = 8 + } + return v + } + // Slow path: consume bits across byte boundaries byte-at-a-time. + // This replaces a bit-by-bit loop (up to 24 getBit() calls) with at most + // ⌈n/8⌉+1 iterations, significantly reducing branch and call overhead. + var result uint32 + for rem := n; rem > 0; { + if br.bytePos >= len(br.data) { + result <<= rem + break + } + take := rem + if take > br.bitPos { + take = br.bitPos + } + br.bitPos -= take + rem -= take + result = (result << take) | ((uint32(br.data[br.bytePos]) >> br.bitPos) & uint32((1<= zgfxHistorySize { + copy(z.history, data[n-zgfxHistorySize:]) + z.historyIdx = 0 + return + } + end := z.historyIdx + n + if end <= zgfxHistorySize { + copy(z.history[z.historyIdx:end], data) + z.historyIdx = end + if z.historyIdx == zgfxHistorySize { + z.historyIdx = 0 + } + } else { + first := zgfxHistorySize - z.historyIdx + copy(z.history[z.historyIdx:], data[:first]) + copy(z.history[:n-first], data[first:]) + z.historyIdx = n - first + } +} + +func (z *zgfxContext) outputLiteral(b byte, out *[]byte) { + z.history[z.historyIdx] = b + z.historyIdx++ + if z.historyIdx == zgfxHistorySize { + z.historyIdx = 0 + } + *out = append(*out, b) +} + +// outputMatch copies `count` bytes starting `distance` bytes back in the +// history into both the output and the history ring. Self-referential +// matches (count > distance) are handled via byte-wise copy after the first +// pass, since the pattern grows as it is written. +func (z *zgfxContext) outputMatch(distance, count int, out *[]byte) { + if distance <= 0 || count <= 0 { + return + } + o := *out + base := len(o) + // Grow output by `count` bytes in a single allocation step. + if cap(o)-base < count { + // Standard slice-growth: at least double cap or fit count, whichever larger. + newCap := max(cap(o)*2, base+count) + grown := make([]byte, base+count, newCap) + copy(grown, o) + o = grown + } else { + o = o[:base+count] + } + + srcIdx := z.historyIdx - distance + if srcIdx < 0 { + srcIdx += zgfxHistorySize + } + + if count <= distance { + // Non-overlapping pattern: 1-2 contiguous copies from history. + end := srcIdx + count + if end <= zgfxHistorySize { + copy(o[base:base+count], z.history[srcIdx:end]) + } else { + first := zgfxHistorySize - srcIdx + copy(o[base:base+first], z.history[srcIdx:]) + copy(o[base+first:base+count], z.history[:count-first]) + } + } else { + // Self-overlapping: copy the initial `distance` bytes from history, + // then expand the pattern in-place within the output buffer (the + // classic LZ77 overlapping copy). + end := srcIdx + distance + if end <= zgfxHistorySize { + copy(o[base:base+distance], z.history[srcIdx:end]) + } else { + first := zgfxHistorySize - srcIdx + copy(o[base:base+first], z.history[srcIdx:]) + copy(o[base+first:base+distance], z.history[:distance-first]) + } + // Overlapping in-place expansion (must be byte-wise; copy() does + // not guarantee overlapping semantics for src==dst+offset). + for i := distance; i < count; i++ { + o[base+i] = o[base+i-distance] + } + } + + *out = o + z.historyWrite(o[base : base+count]) +} + +// Decompress decompresses a ZGFX compressed segment payload. +// The payload must NOT include the 1-byte segment header (flags byte). +// In RDP8 ZGFX, the last byte of the payload encodes the number of +// padding bits to subtract from the total bit count. +// +// buf is an optional caller-supplied backing buffer (e.g. from a sync.Pool). +// The returned slice may use buf's backing array or a new one if buf was too +// small; callers that pool the output must use the RETURNED slice, not buf. +func (z *zgfxContext) Decompress(data []byte, buf []byte) []byte { + if len(data) < 2 { + // Need at least 1 byte of compressed data + 1 byte of padding count + return nil + } + + br := newBitReaderWithCount(data) + // Use the caller-supplied buffer if it has enough capacity; otherwise fall + // back to a fresh allocation sized at 3× the compressed input. + estSize := len(data) * 3 + var out []byte + if cap(buf) >= estSize { + out = buf[:0] + } else { + out = make([]byte, 0, estSize) + } + + for br.hasBitsRemaining() { + token, ok := z.decodeToken(br) + if !ok { + break + } + // decodeToken already consumed prefixLen bits. + + if token.tokenType == tokenLiteral { + if br.bitsRemaining < uint32(token.valueBits) { + break + } + value := token.valueBase + br.getBits(token.valueBits) + br.bitsRemaining -= uint32(token.valueBits) + z.outputLiteral(byte(value), &out) + } else { + // Match token + if br.bitsRemaining < uint32(token.valueBits) { + break + } + distance := int(token.valueBase + br.getBits(token.valueBits)) + br.bitsRemaining -= uint32(token.valueBits) + + if distance != 0 { + // Match: copy from history + count := z.decodeMatchCount(br) + z.outputMatch(distance, count, &out) + } else { + // Unencoded: read raw bytes + if br.bitsRemaining < 15 { + break + } + rawCount := int(br.getBits(15)) + br.bitsRemaining -= 15 + // Discard remaining bits in current byte to align to byte boundary + // (equivalent to FreeRDP's cBitsCurrent = 0; BitsCurrent = 0;) + if br.bitPos < 8 { + br.bitsRemaining -= uint32(br.bitPos) + br.bytePos++ + br.bitPos = 8 + } + if br.bytePos+rawCount > len(br.data) || uint32(rawCount)*8 > br.bitsRemaining { + break + } + rawBytes := br.data[br.bytePos : br.bytePos+rawCount] + br.bytePos += rawCount + br.bitsRemaining -= uint32(rawCount) * 8 + z.historyWrite(rawBytes) + out = append(out, rawBytes...) + } + } + } + + return out +} + +// decodeToken consumes the next Huffman prefix (1..9 bits) from the input +// stream using the precomputed zgfxTokenLut. Returns the decoded token and +// true on success; false if the remaining input is too short or malformed. +// +// Unlike the old implementation, this version no longer uses a per-bit +// linear search through the 38-entry token table — the dominant ZGFX cost +// during real RDPGFX traffic. +func (z *zgfxContext) decodeToken(br *bitReader) (zgfxToken, bool) { + val, avail := br.peek9() + if avail == 0 { + return zgfxToken{}, false + } + e := zgfxTokenLut[val] + if e.prefixLen == 0 || uint8(avail) < e.prefixLen { + return zgfxToken{}, false + } + br.consumeBits(e.prefixLen) + return zgfxToken{ + prefixLen: e.prefixLen, + valueBits: e.valueBits, + tokenType: e.tokenType, + valueBase: e.valueBase, + }, true +} + +// decodeMatchCount decodes the match length using FreeRDP's algorithm: +// +// 0 → 3 +// 10 + 2 bits → 4 + value (4..7) +// 110 + 3 bits → 8 + value (8..15) +// 1110 + 4 bits → 16 + value (16..31) +// ... and so on (each additional leading 1 doubles the base and adds 1 extra bit) +func (z *zgfxContext) decodeMatchCount(br *bitReader) int { + bit := br.getBit() + br.bitsRemaining-- + if bit == 0 { + return 3 + } + + count := 4 + extra := uint8(2) + + bit = br.getBit() + br.bitsRemaining-- + for bit == 1 { + count <<= 1 + extra++ + bit = br.getBit() + br.bitsRemaining-- + } + + if br.bitsRemaining < uint32(extra) { + return count + } + count += int(br.getBits(extra)) + br.bitsRemaining -= uint32(extra) + return count +} diff --git a/plugin/rdpsnd/aac/decoder.go b/plugin/rdpsnd/aac/decoder.go new file mode 100644 index 0000000..3cacc2b --- /dev/null +++ b/plugin/rdpsnd/aac/decoder.go @@ -0,0 +1,18 @@ +// Package aac provides AAC-to-PCM decoding for RDPSND audio streams. +package aac + +import "git.zeroonesoft.cn/golib/rdplib/plugin/rdpsnd" + +// Decoder decodes MPEG-4 AAC packets into signed 16-bit little-endian PCM. +type Decoder interface { + // Decode decodes one raw AAC packet. Returns nil, nil for empty input. + Decode(data []byte) ([]byte, error) + // Close releases all resources held by the decoder. + Close() +} + +// New creates a Decoder for the given RDPSND AudioFormat. +// Returns an error on platforms where AAC decoding is not supported. +func New(format rdpsnd.AudioFormat) (Decoder, error) { + return newDecoder(format) +} diff --git a/plugin/rdpsnd/aac/decoder_darwin.go b/plugin/rdpsnd/aac/decoder_darwin.go new file mode 100644 index 0000000..7499b3f --- /dev/null +++ b/plugin/rdpsnd/aac/decoder_darwin.go @@ -0,0 +1,210 @@ +//go:build darwin && cgo + +package aac + +/* +#cgo LDFLAGS: -framework AudioToolbox +#cgo nocallback grdp_aac_converter_new +#cgo nocallback grdp_aac_decode +#cgo nocallback AudioConverterDispose +#cgo noescape grdp_aac_converter_new +#cgo noescape grdp_aac_decode +#include +#include + +typedef struct { + const uint8_t *data; + uint32_t size; + int consumed; +} grdp_aac_input_t; + +// AudioConverter input callback — called by AudioToolbox to pull one AAC +// packet per invocation. Returns noErr on the first call, then -1 (no data) +// on subsequent calls to signal end-of-input for the current decode round. +static OSStatus grdp_aac_input_proc( + AudioConverterRef inConverter, + UInt32 *ioNumberDataPackets, + AudioBufferList *ioData, + AudioStreamPacketDescription **outDataPacketDescription, + void *inUserData) +{ + grdp_aac_input_t *state = (grdp_aac_input_t *)inUserData; + if (state->consumed || state->size == 0) { + *ioNumberDataPackets = 0; + ioData->mBuffers[0].mData = NULL; + ioData->mBuffers[0].mDataByteSize = 0; + return -1; + } + *ioNumberDataPackets = 1; + ioData->mBuffers[0].mData = (void *)state->data; + ioData->mBuffers[0].mDataByteSize = state->size; + if (outDataPacketDescription != NULL) { + static AudioStreamPacketDescription desc; + desc.mStartOffset = 0; + desc.mDataByteSize = state->size; + desc.mVariableFramesInPacket = 0; + *outDataPacketDescription = &desc; + } + state->consumed = 1; + return noErr; +} + +// grdp_aac_converter_new creates an AudioConverter that decodes MPEG-4 AAC at +// the given sample rate / channel count. asc/ascLen is the AudioSpecificConfig +// stored in the WAVEFORMATEX ExtraData field; it may be NULL/0 for ADTS input. +static AudioConverterRef grdp_aac_converter_new( + double sampleRate, + uint32_t channels, + const uint8_t *asc, + uint32_t ascLen, + OSStatus *outErr) +{ + AudioStreamBasicDescription inFmt = { + .mSampleRate = sampleRate, + .mFormatID = kAudioFormatMPEG4AAC, + .mChannelsPerFrame = channels, + }; + AudioStreamBasicDescription outFmt = { + .mSampleRate = sampleRate, + .mFormatID = kAudioFormatLinearPCM, + .mFormatFlags = kLinearPCMFormatFlagIsSignedInteger | + kLinearPCMFormatFlagIsPacked, + .mBitsPerChannel = 16, + .mChannelsPerFrame = channels, + .mBytesPerFrame = (uint32_t)(2 * channels), + .mFramesPerPacket = 1, + .mBytesPerPacket = (uint32_t)(2 * channels), + }; + AudioConverterRef conv = NULL; + OSStatus err = AudioConverterNew(&inFmt, &outFmt, &conv); + if (err != noErr) { *outErr = err; return NULL; } + if (ascLen > 0 && asc != NULL) { + err = AudioConverterSetProperty( + conv, + kAudioConverterDecompressionMagicCookie, + ascLen, asc); + if (err != noErr) { + AudioConverterDispose(conv); + *outErr = err; + return NULL; + } + } + *outErr = noErr; + return conv; +} + +// grdp_aac_decode decodes one AAC packet (inData/inSize) into signed 16-bit +// little-endian PCM stored in outBuf. *outBytesWritten is set to the number +// of bytes actually written. Returns noErr on success. +static OSStatus grdp_aac_decode( + AudioConverterRef conv, + const uint8_t *inData, + uint32_t inSize, + uint8_t *outBuf, + uint32_t outBufSize, + uint32_t *outBytesWritten) +{ + grdp_aac_input_t state = { inData, inSize, 0 }; + AudioBufferList outList = { + .mNumberBuffers = 1, + .mBuffers[0] = { + .mNumberChannels = 0, + .mDataByteSize = outBufSize, + .mData = outBuf, + }, + }; + // AAC-LC produces 1024 samples/channel per frame; AAC-HE produces 2048. + // 4096 is large enough for both at up to 2 channels. + UInt32 numFrames = 4096; + OSStatus err = AudioConverterFillComplexBuffer( + conv, grdp_aac_input_proc, &state, &numFrames, &outList, NULL); + // -1 from the input callback signals "no more data"; that is not an error. + if (err != noErr && err != -1) { + *outBytesWritten = 0; + return err; + } + *outBytesWritten = outList.mBuffers[0].mDataByteSize; + return noErr; +} +*/ +import "C" + +import ( + "fmt" + "log/slog" + "runtime" + "unsafe" + + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpsnd" +) + +// darwinDecoder decodes MPEG-4 AAC audio into signed 16-bit PCM using +// macOS AudioToolbox (hardware-accelerated on Apple Silicon / Intel iGPU). +type darwinDecoder struct { + conv C.AudioConverterRef + channels int + outBuf []byte // reusable output scratch buffer +} + +func newDecoder(format rdpsnd.AudioFormat) (Decoder, error) { + var oscErr C.OSStatus + var ascPtr *C.uint8_t + ascLen := C.uint32_t(0) + if len(format.ExtraData) > 0 { + ascPtr = (*C.uint8_t)(unsafe.Pointer(&format.ExtraData[0])) + ascLen = C.uint32_t(len(format.ExtraData)) + } + conv := C.grdp_aac_converter_new( + C.double(format.SamplesPerSec), + C.uint32_t(format.Channels), + ascPtr, ascLen, + &oscErr) + if len(format.ExtraData) > 0 { + runtime.KeepAlive(format.ExtraData) + } + if conv == nil { + return nil, fmt.Errorf("AudioToolbox: AudioConverterNew failed (err=%d)", int32(oscErr)) + } + slog.Info("AAC decoder created", + "rate", format.SamplesPerSec, + "channels", format.Channels, + "ascLen", len(format.ExtraData)) + // 4096 samples × channels × 2 bytes/sample — large enough for AAC-HE at 2 ch. + outBufSize := 4096 * int(format.Channels) * 2 + return &darwinDecoder{conv: conv, channels: int(format.Channels), outBuf: make([]byte, outBufSize)}, nil +} + +// Decode decodes one raw AAC packet into signed 16-bit little-endian PCM. +// Returns nil, nil when data is empty. +func (d *darwinDecoder) Decode(data []byte) ([]byte, error) { + if len(data) == 0 { + return nil, nil + } + var written C.uint32_t + err := C.grdp_aac_decode( + d.conv, + (*C.uint8_t)(unsafe.Pointer(&data[0])), + C.uint32_t(len(data)), + (*C.uint8_t)(unsafe.Pointer(&d.outBuf[0])), + C.uint32_t(len(d.outBuf)), + &written) + runtime.KeepAlive(data) + runtime.KeepAlive(d.outBuf) + if err != 0 { + return nil, fmt.Errorf("AudioToolbox: decode error %d", int32(err)) + } + if written == 0 { + return nil, nil + } + out := make([]byte, int(written)) + copy(out, d.outBuf[:written]) + return out, nil +} + +// Close releases the AudioConverter. +func (d *darwinDecoder) Close() { + if d.conv != nil { + C.AudioConverterDispose(d.conv) + d.conv = nil + } +} diff --git a/plugin/rdpsnd/aac/decoder_stub.go b/plugin/rdpsnd/aac/decoder_stub.go new file mode 100644 index 0000000..df575a0 --- /dev/null +++ b/plugin/rdpsnd/aac/decoder_stub.go @@ -0,0 +1,22 @@ +//go:build !darwin || !cgo + +package aac + +import ( + "errors" + + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpsnd" +) + +// stubDecoder is used on platforms without AudioToolbox support. +type stubDecoder struct{} + +func newDecoder(_ rdpsnd.AudioFormat) (Decoder, error) { + return nil, errors.New("AAC decoding not supported on this platform") +} + +func (d *stubDecoder) Decode(_ []byte) ([]byte, error) { + return nil, errors.New("AAC decoding not supported on this platform") +} + +func (d *stubDecoder) Close() {} diff --git a/plugin/rdpsnd/rdpsnd.go b/plugin/rdpsnd/rdpsnd.go new file mode 100644 index 0000000..fd1967a --- /dev/null +++ b/plugin/rdpsnd/rdpsnd.go @@ -0,0 +1,524 @@ +// Package rdpsnd implements the RDPSND (Audio Output Virtual Channel Extension) +// protocol (MS-RDPEA) for server-to-client audio redirection. +// +// It can operate over either a static virtual channel ("rdpsnd") or +// a dynamic virtual channel (AUDIO_PLAYBACK_DVC / AUDIO_PLAYBACK_LOSSY_DVC). +package rdpsnd + +import ( + "bytes" + "encoding/binary" + "fmt" + "log/slog" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/plugin" +) + +const ( + ChannelName = plugin.RDPSND_SVC_CHANNEL_NAME + ChannelOption = plugin.CHANNEL_OPTION_INITIALIZED | + plugin.CHANNEL_OPTION_ENCRYPT_RDP +) + +// RDPSND PDU types (MS-RDPEA 2.2) +const ( + SNDC_CLOSE = 0x01 + SNDC_WAVE = 0x02 + SNDC_SETVOLUME = 0x03 + SNDC_SETPITCH = 0x04 + SNDC_WAVECONFIRM = 0x05 + SNDC_TRAINING = 0x06 + SNDC_FORMATS = 0x07 + SNDC_CRYPTKEY = 0x08 + SNDC_WAVEENCRYPT = 0x09 + SNDC_UDPWAVE = 0x0A + SNDC_UDPWAVELAST = 0x0B + SNDC_QUALITYMODE = 0x0C + SNDC_WAVE2 = 0x0D +) + +// RDPSND capabilities flags +const ( + TSSNDCAPS_ALIVE = 0x00000001 + TSSNDCAPS_VOLUME = 0x00000002 + TSSNDCAPS_PITCH = 0x00000004 +) + +// Quality mode values (MS-RDPEA 2.2.2.9) +const ( + DYNAMIC_QUALITY = 0x0000 + MEDIUM_QUALITY = 0x0002 + HIGH_QUALITY = 0x0001 +) + +// Audio format tags +const ( + WAVE_FORMAT_PCM = 0x0001 + WAVE_FORMAT_ADPCM = 0x0002 + WAVE_FORMAT_ALAW = 0x0006 + WAVE_FORMAT_MULAW = 0x0007 + WAVE_FORMAT_AAC = 0x00FF // MPEG-4 AAC (AudioSpecificConfig in ExtraData) +) + +// RDPSND version +// gnome-remote-desktop (grd-rdp-dvc-audio-playback.c) requires +// clientVersion >= 8 (CHANNEL_VERSION_WIN_8). FreeRDP WIN_7=6, WIN_8=8. +const ( + RDPSND_VERSION_MAJOR = 0x08 +) + +// AudioFormat represents a WAVEFORMATEX structure. +type AudioFormat struct { + Tag uint16 + Channels uint16 + SamplesPerSec uint32 + AvgBytesPerSec uint32 + BlockAlign uint16 + BitsPerSample uint16 + ExtraData []byte +} + +func (f AudioFormat) String() string { + var name string + switch f.Tag { + case WAVE_FORMAT_PCM: + name = "PCM" + case WAVE_FORMAT_ADPCM: + name = "ADPCM" + case WAVE_FORMAT_ALAW: + name = "A-Law" + case WAVE_FORMAT_MULAW: + name = "μ-Law" + case WAVE_FORMAT_AAC: + name = "AAC" + default: + name = fmt.Sprintf("0x%04x", f.Tag) + } + return fmt.Sprintf("%s %dHz %dch %dbit", name, f.SamplesPerSec, f.Channels, f.BitsPerSample) +} + +func (f AudioFormat) IsPCM() bool { + return f.Tag == WAVE_FORMAT_PCM +} + +// IsAAC reports whether the format uses MPEG-4 AAC encoding. +func (f AudioFormat) IsAAC() bool { + return f.Tag == WAVE_FORMAT_AAC +} + +func (f AudioFormat) pack() []byte { + b := make([]byte, 18+len(f.ExtraData)) + binary.LittleEndian.PutUint16(b[0:], f.Tag) + binary.LittleEndian.PutUint16(b[2:], f.Channels) + binary.LittleEndian.PutUint32(b[4:], f.SamplesPerSec) + binary.LittleEndian.PutUint32(b[8:], f.AvgBytesPerSec) + binary.LittleEndian.PutUint16(b[12:], f.BlockAlign) + binary.LittleEndian.PutUint16(b[14:], f.BitsPerSample) + binary.LittleEndian.PutUint16(b[16:], uint16(len(f.ExtraData))) + copy(b[18:], f.ExtraData) + return b +} + +func unpackAudioFormat(data []byte, offset int) (AudioFormat, int) { + if len(data)-offset < 18 { + return AudioFormat{}, offset + } + f := AudioFormat{ + Tag: binary.LittleEndian.Uint16(data[offset:]), + Channels: binary.LittleEndian.Uint16(data[offset+2:]), + SamplesPerSec: binary.LittleEndian.Uint32(data[offset+4:]), + AvgBytesPerSec: binary.LittleEndian.Uint32(data[offset+8:]), + BlockAlign: binary.LittleEndian.Uint16(data[offset+12:]), + BitsPerSample: binary.LittleEndian.Uint16(data[offset+14:]), + } + cbSize := int(binary.LittleEndian.Uint16(data[offset+16:])) + if offset+18+cbSize <= len(data) { + f.ExtraData = make([]byte, cbSize) + copy(f.ExtraData, data[offset+18:offset+18+cbSize]) + } + return f, offset + 18 + cbSize +} + +// Handler implements the RDPSND protocol over a static virtual channel. +// It also serves as the DVC audio handler via ProcessData. +type Handler struct { + channelSender core.ChannelSender + + serverFormats []AudioFormat + clientFormatIndices []int + activeFormatIndex int + + // Wave state + waveTimestamp uint16 + waveBlockNo uint8 + pendingWaveHdr [4]byte // backing array for pendingWave initial bytes (avoids alloc) + pendingWave []byte + expectingWave bool + + // DVC send callback for the current message's channel + dvcSendFunc func([]byte) + + // viaDvc tracks whether the current message arrived via DVC + viaDvc bool + + // Application callback: called with the active AudioFormat and PCM data + onAudio func(AudioFormat, []byte) + + // onAudioReset is called when the server closes the audio channel + // (SNDC_CLOSE). The application should flush its audio playback buffer + // so that stale audio from before a seek does not keep playing. + onAudioReset func() + + // muted 为 true 时丢弃 wave 数据不回调 onAudio,但协议握手与 + // wave 确认照常——服务器认为音频已被重定向而保持静音 + //(对应 mstsc「不播放」模式)。 + muted bool +} + +// NewHandler creates a new RDPSND handler. +// onAudio is called with the active AudioFormat and PCM audio data for each wave. +func NewHandler(onAudio func(AudioFormat, []byte)) *Handler { + return &Handler{ + activeFormatIndex: -1, + onAudio: onAudio, + } +} + +// SetAudioResetCallback sets a function that is called when the server +// closes the audio channel (e.g. on media seek). The application should +// flush its audio playback buffer in this callback. +func (h *Handler) SetAudioResetCallback(f func()) { + h.onAudioReset = f +} + +// SetMuted controls wave playback: muted=true 丢弃音频数据(不回调 +// onAudio),协议层格式协商与 wave 确认照常进行。 +func (h *Handler) SetMuted(m bool) { + h.muted = m +} + +// --- plugin.ChannelTransport interface --- + +func (h *Handler) GetType() (string, uint32) { + return ChannelName, ChannelOption +} + +func (h *Handler) Sender(s core.ChannelSender) { + h.channelSender = s +} + +// Process handles data from the static virtual channel (already reassembled). +func (h *Handler) Process(s []byte) { + defer func() { + if r := recover(); r != nil { + slog.Error("rdpsnd: panic in Process", "err", r) + } + }() + h.viaDvc = false + h.ProcessData(s) +} + +// ProcessData processes a reassembled RDPSND PDU payload. +// This is used by both the static VChannel path and the DVC path. +func (h *Handler) ProcessData(data []byte) { + if h.expectingWave { + h.processWaveBody(data) + return + } + + if len(data) < 4 { + return + } + + msgType := data[0] + // data[1] is bPad + bodySize := int(binary.LittleEndian.Uint16(data[2:4])) + body := data[4:] + if bodySize < len(body) { + body = body[:bodySize] + } + + switch msgType { + case SNDC_FORMATS: + h.processServerFormats(body) + case SNDC_TRAINING: + h.processTraining(body) + case SNDC_WAVE: + h.processWaveInfo(body) + case SNDC_WAVE2: + h.processWave2(body) + case SNDC_CLOSE: + slog.Debug("rdpsnd: server closed audio channel") + if h.onAudioReset != nil { + h.onAudioReset() + } + case SNDC_SETVOLUME, SNDC_QUALITYMODE: + // ignored + default: + slog.Debug("rdpsnd: unknown msgType", "type", fmt.Sprintf("0x%02x", msgType)) + } +} + +// --- Server Audio Formats and Version (MS-RDPEA 2.2.2.1) --- + +func (h *Handler) processServerFormats(body []byte) { + if len(body) < 20 { + slog.Warn("rdpsnd: Server Formats PDU too short") + return + } + + dwFlags := binary.LittleEndian.Uint32(body[0:]) + _ = dwFlags + wNumberOfFormats := binary.LittleEndian.Uint16(body[14:]) + wVersion := binary.LittleEndian.Uint16(body[17:]) + + slog.Debug("rdpsnd: Server Formats", "version", wVersion, "numFormats", wNumberOfFormats) + + offset := 20 + h.serverFormats = nil + for i := 0; i < int(wNumberOfFormats); i++ { + fmt, newOffset := unpackAudioFormat(body, offset) + if newOffset == offset { + break + } + h.serverFormats = append(h.serverFormats, fmt) + slog.Debug("rdpsnd: server format", "idx", i, "fmt", fmt) + offset = newOffset + } + + // Prefer AAC formats first (hardware-decoded on macOS), then fall back to PCM. + h.clientFormatIndices = nil + for i, f := range h.serverFormats { + if f.IsAAC() { + h.clientFormatIndices = append(h.clientFormatIndices, i) + } + } + for i, f := range h.serverFormats { + if f.IsPCM() && (f.BitsPerSample == 8 || f.BitsPerSample == 16) && (f.Channels == 1 || f.Channels == 2) { + h.clientFormatIndices = append(h.clientFormatIndices, i) + } + } + + if len(h.clientFormatIndices) == 0 { + slog.Warn("rdpsnd: no supported PCM format found") + } + + h.sendClientFormats(wVersion) +} + +func (h *Handler) sendClientFormats(serverVersion uint16) { + version := min(serverVersion, RDPSND_VERSION_MAJOR) + + formatData := &bytes.Buffer{} + for _, idx := range h.clientFormatIndices { + formatData.Write(h.serverFormats[idx].pack()) + } + + // Header: dwFlags(4) + dwVolume(4) + dwPitch(4) + wDGramPort(2) + // + wNumberOfFormats(2) + cLastBlockConfirmed(1) + wVersion(2) + bPad(1) + hdr := &bytes.Buffer{} + binary.Write(hdr, binary.LittleEndian, uint32(TSSNDCAPS_ALIVE)) // dwFlags + binary.Write(hdr, binary.LittleEndian, uint32(0)) // dwVolume + binary.Write(hdr, binary.LittleEndian, uint32(0)) // dwPitch + binary.Write(hdr, binary.LittleEndian, uint16(0)) // wDGramPort + binary.Write(hdr, binary.LittleEndian, uint16(len(h.clientFormatIndices))) + hdr.WriteByte(0) // cLastBlockConfirmed + binary.Write(hdr, binary.LittleEndian, version) // wVersion + hdr.WriteByte(0) // bPad + + body := append(hdr.Bytes(), formatData.Bytes()...) + + pdu := &bytes.Buffer{} + pdu.WriteByte(SNDC_FORMATS) // msgType + pdu.WriteByte(0) // bPad + binary.Write(pdu, binary.LittleEndian, uint16(len(body))) + pdu.Write(body) + + h.send(pdu.Bytes()) + slog.Debug("rdpsnd: sent Client Formats", "version", version, "numFormats", len(h.clientFormatIndices)) + + // FreeRDP sends a Quality Mode PDU immediately after Client Formats. + // Without it, Windows waits (up to ~10 seconds) before sending Training. + h.sendQualityMode() +} + +// --- Quality Mode (MS-RDPEA 2.2.2.9) --- + +func (h *Handler) sendQualityMode() { + pdu := [8]byte{ + SNDC_QUALITYMODE, 0, + 4, 0, // bodySize = 4 (little-endian uint16) + } + binary.LittleEndian.PutUint16(pdu[4:], HIGH_QUALITY) + // pdu[6:8] = Reserved, already zero + h.send(pdu[:]) + slog.Debug("rdpsnd: sent QualityMode") +} + +// --- Training (MS-RDPEA 2.2.2.3) --- + +func (h *Handler) processTraining(body []byte) { + if len(body) < 4 { + return + } + wTimeStamp := binary.LittleEndian.Uint16(body[0:]) + wPackSize := binary.LittleEndian.Uint16(body[2:]) + slog.Debug("rdpsnd: Training", "timestamp", wTimeStamp, "packSize", wPackSize) + + pdu := [8]byte{SNDC_TRAINING, 0, 4, 0} // msgType, bPad, bodySize=4 (LE) + binary.LittleEndian.PutUint16(pdu[4:], wTimeStamp) + binary.LittleEndian.PutUint16(pdu[6:], wPackSize) + h.send(pdu[:]) + slog.Debug("rdpsnd: sent Training Confirm") +} + +// --- Wave Info / Wave Data (MS-RDPEA 2.2.2.5 / 2.2.2.6) --- + +func (h *Handler) processWaveInfo(body []byte) { + if len(body) < 12 { + slog.Warn("rdpsnd: WaveInfo body too short") + return + } + + wTimeStamp := binary.LittleEndian.Uint16(body[0:]) + wFormatNo := binary.LittleEndian.Uint16(body[2:]) + cBlockNo := body[4] + copy(h.pendingWaveHdr[:], body[8:12]) + + h.waveTimestamp = wTimeStamp + h.waveBlockNo = cBlockNo + h.pendingWave = h.pendingWaveHdr[:] + + if int(wFormatNo) < len(h.clientFormatIndices) { + serverIdx := h.clientFormatIndices[wFormatNo] + h.activeFormatIndex = serverIdx + } else { + slog.Warn("rdpsnd: WaveInfo format index out of range", "idx", wFormatNo, "max", len(h.clientFormatIndices)) + } + + h.expectingWave = true + slog.Debug("rdpsnd: WaveInfo", "ts", wTimeStamp, "fmt", wFormatNo, "block", cBlockNo) +} + +func (h *Handler) processWaveBody(data []byte) { + h.expectingWave = false + // First 4 bytes are padding (duplicate of WaveInfo header) + var audioData []byte + if len(data) > 4 { + audioData = append(h.pendingWave, data[4:]...) + } else { + audioData = h.pendingWave + } + h.pendingWave = nil + + slog.Debug("rdpsnd: Wave data", "len", len(audioData)) + h.deliverAudio(audioData) + var confirmFmt AudioFormat + if h.activeFormatIndex >= 0 && h.activeFormatIndex < len(h.serverFormats) { + confirmFmt = h.serverFormats[h.activeFormatIndex] + } + h.sendWaveConfirm(waveConfirmTimestamp(h.waveTimestamp, audioData, confirmFmt), h.waveBlockNo) +} + +// --- Wave2 (MS-RDPEA 2.2.2.7) --- + +func (h *Handler) processWave2(body []byte) { + if len(body) < 12 { + slog.Warn("rdpsnd: Wave2 body too short") + return + } + + wTimeStamp := binary.LittleEndian.Uint16(body[0:]) + wFormatNo := binary.LittleEndian.Uint16(body[2:]) + cBlockNo := body[4] + audioData := body[12:] + + if int(wFormatNo) < len(h.clientFormatIndices) { + serverIdx := h.clientFormatIndices[wFormatNo] + h.activeFormatIndex = serverIdx + } else { + slog.Warn("rdpsnd: Wave2 format index out of range", "idx", wFormatNo, "max", len(h.clientFormatIndices)) + } + + slog.Debug("rdpsnd: Wave2", "ts", wTimeStamp, "fmt", wFormatNo, "block", cBlockNo, "dataLen", len(audioData)) + h.deliverAudio(audioData) + var confirmFmt AudioFormat + if h.activeFormatIndex >= 0 && h.activeFormatIndex < len(h.serverFormats) { + confirmFmt = h.serverFormats[h.activeFormatIndex] + } + h.sendWaveConfirm(waveConfirmTimestamp(wTimeStamp, audioData, confirmFmt), cBlockNo) +} + +// --- Wave Confirm (MS-RDPEA 2.2.2.8) --- + +// waveConfirmTimestamp computes the wTimeStamp for WAVE_CONFIRM_PDU. +// MS-RDPEA §2.2.2.8: the confirmed timestamp MUST be the server's timestamp +// PLUS the estimated playback duration of the audio data in milliseconds. +// This allows the server to pace audio delivery accurately. +func waveConfirmTimestamp(serverTs uint16, audioData []byte, fmt AudioFormat) uint16 { + if fmt.AvgBytesPerSec == 0 { + return serverTs + } + playMs := uint32(len(audioData)) * 1000 / fmt.AvgBytesPerSec + return serverTs + uint16(playMs) +} + +func (h *Handler) sendWaveConfirm(timestamp uint16, blockNo uint8) { + var pdu [8]byte + pdu[0] = SNDC_WAVECONFIRM + // pdu[1] = bPad (zero) + pdu[2] = 4 // bodySize = 4 (little-endian uint16, high byte stays 0) + binary.LittleEndian.PutUint16(pdu[4:], timestamp) + pdu[6] = blockNo + // pdu[7] = bPad (zero) + h.send(pdu[:]) + slog.Debug("rdpsnd: sent WaveConfirm", "ts", timestamp, "block", blockNo) +} + +// --- Audio delivery --- + +func (h *Handler) deliverAudio(data []byte) { + if h.muted || h.onAudio == nil || h.activeFormatIndex < 0 || h.activeFormatIndex >= len(h.serverFormats) { + return + } + h.onAudio(h.serverFormats[h.activeFormatIndex], data) +} + +// --- Send helpers --- + +// send sends a response on the same path that the current message arrived on. +// Static channel messages get static channel responses; DVC messages get DVC responses. +func (h *Handler) send(data []byte) { + if h.viaDvc && h.dvcSendFunc != nil { + h.dvcSendFunc(data) + } else if h.channelSender != nil { + h.channelSender.SendToChannel(ChannelName, data) + } +} + +// --- DVC adapter --- + +// DvcAdapter wraps an rdpsnd Handler to work as a DVC channel handler. +// Each DVC channel gets its own adapter so responses go to the correct channel. +type DvcAdapter struct { + handler *Handler + sendFunc func([]byte) +} + +// NewDvcAdapter creates a DVC adapter that routes audio data to the given Handler. +func NewDvcAdapter(handler *Handler) *DvcAdapter { + return &DvcAdapter{handler: handler} +} + +// Process implements drdynvc.DvcChannelHandler. +func (a *DvcAdapter) Process(data []byte) { + a.handler.viaDvc = true + a.handler.dvcSendFunc = a.sendFunc + a.handler.ProcessData(data) +} + +// SetSendFunc is called by the DVC client to provide the send function. +func (a *DvcAdapter) SetSendFunc(fn func([]byte)) { + a.sendFunc = fn +} diff --git a/protocol/lic/lic.go b/protocol/lic/lic.go new file mode 100644 index 0000000..132057f --- /dev/null +++ b/protocol/lic/lic.go @@ -0,0 +1,183 @@ +package lic + +import ( + "io" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +const ( + LICENSE_REQUEST = 0x01 + PLATFORM_CHALLENGE = 0x02 + NEW_LICENSE = 0x03 + UPGRADE_LICENSE = 0x04 + LICENSE_INFO = 0x12 + NEW_LICENSE_REQUEST = 0x13 + PLATFORM_CHALLENGE_RESPONSE = 0x15 + ERROR_ALERT = 0xFF +) + +// error code +const ( + ERR_INVALID_SERVER_CERTIFICATE = 0x00000001 + ERR_NO_LICENSE = 0x00000002 + ERR_INVALID_SCOPE = 0x00000004 + ERR_NO_LICENSE_SERVER = 0x00000006 + STATUS_VALID_CLIENT = 0x00000007 + ERR_INVALID_CLIENT = 0x00000008 + ERR_INVALID_PRODUCTID = 0x0000000B + ERR_INVALID_MESSAGE_LEN = 0x0000000C + ERR_INVALID_MAC = 0x00000003 +) + +// state transition +const ( + ST_TOTAL_ABORT = 0x00000001 + ST_NO_TRANSITION = 0x00000002 + ST_RESET_PHASE_TO_START = 0x00000003 + ST_RESEND_LAST_MESSAGE = 0x00000004 +) + +/* +""" +@summary: Binary blob data type +@see: http://msdn.microsoft.com/en-us/library/cc240481.aspx +""" +*/ +type BinaryBlobType uint16 + +const ( + BB_ANY_BLOB = 0x0000 + BB_DATA_BLOB = 0x0001 + BB_RANDOM_BLOB = 0x0002 + BB_CERTIFICATE_BLOB = 0x0003 + BB_ERROR_BLOB = 0x0004 + BB_ENCRYPTED_DATA_BLOB = 0x0009 + BB_KEY_EXCHG_ALG_BLOB = 0x000D + BB_SCOPE_BLOB = 0x000E + BB_CLIENT_USER_NAME_BLOB = 0x000F + BB_CLIENT_MACHINE_NAME_BLOB = 0x0010 +) + +type ErrorMessage struct { + DwErrorCode uint32 + DwStateTransaction uint32 + Blob []byte +} + +func readErrorMessage(r io.Reader) *ErrorMessage { + m := &ErrorMessage{} + m.DwErrorCode, _ = core.ReadUInt32LE(r) + m.DwStateTransaction, _ = core.ReadUInt32LE(r) + return m +} + +type LicensePacket struct { + BMsgtype uint8 + Flag uint8 + WMsgSize uint16 + LicensingMessage any +} + +func ReadLicensePacket(r io.Reader) *LicensePacket { + l := &LicensePacket{} + l.BMsgtype, _ = core.ReadUInt8(r) + l.Flag, _ = core.ReadUInt8(r) + l.WMsgSize, _ = core.ReadUint16LE(r) + + switch l.BMsgtype { + case ERROR_ALERT: + l.LicensingMessage = readErrorMessage(r) + default: + l.LicensingMessage, _ = core.ReadBytes(int(l.WMsgSize-4), r) + } + + return l +} + +/* +""" +@summary: Blob use by license manager to exchange security data +@see: http://msdn.microsoft.com/en-us/library/cc240481.aspx +""" +*/ +type LicenseBinaryBlob struct { + WBlobType uint16 `struc:"little"` + WBlobLen uint16 `struc:"little"` + BlobData []byte `struc:"sizefrom=WBlobLen"` +} + +func NewLicenseBinaryBlob(WBlobType uint16) *LicenseBinaryBlob { + return &LicenseBinaryBlob{} +} + +/* +""" +@summary: License server product information +@see: http://msdn.microsoft.com/en-us/library/cc241915.aspx +""" +*/ +type ProductInformation struct { + DwVersion uint32 `struc:"little"` + CbCompanyName uint32 `struc:"little"` + //may contain "Microsoft Corporation" from server microsoft + PbCompanyName []byte `struc:"sizefrom=CbCompanyName"` + CbProductId uint32 `struc:"little"` + //may contain "A02" from microsoft license server + PbProductId []byte `struc:"sizefrom=CbProductId"` +} + +/* +@summary: Send by server to signal license request + + server -> client + +@see: http://msdn.microsoft.com/en-us/library/cc241914.aspx +*/ +type ServerLicenseRequest struct { + ServerRandom []byte `struc:"[32]byte"` + ProductInfo ProductInformation `struc:"little"` + KeyExchangeList LicenseBinaryBlob `struc:"little"` + ServerCertificate LicenseBinaryBlob `struc:"little"` + //ScopeList ScopeList +} + +/* +@summary: Send by client to ask new license for client. + RDPY doesn'support license reuse, need it in futur version +@see: http://msdn.microsoft.com/en-us/library/cc241918.aspx + #RSA and must be only RSA + #pure microsoft client ;-) + #http://msdn.microsoft.com/en-us/library/1040af38-c733-4fb3-acd1-8db8cc979eda#id10 +*/ + +type ClientNewLicenseRequest struct { + PreferredKeyExchangeAlg uint32 `struc:"little"` + PlatformId uint32 `struc:"little"` + ClientRandom []byte `struc:"[32]byte"` + EncryptedPreMasterSecret LicenseBinaryBlob `struc:"little"` + ClientUserName LicenseBinaryBlob `struc:"little"` + ClientMachineName LicenseBinaryBlob `struc:"little"` +} + +/* +@summary: challenge send from server to client +@see: http://msdn.microsoft.com/en-us/library/cc241921.aspx +*/ +type ServerPlatformChallenge struct { + ConnectFlags uint32 + EncryptedPlatformChallenge LicenseBinaryBlob + MACData [16]byte +} + +/* +""" +@summary: client challenge response +@see: http://msdn.microsoft.com/en-us/library/cc241922.aspx +""" +*/ +type ClientPLatformChallengeResponse struct { + EncryptedPlatformChallengeResponse LicenseBinaryBlob + EncryptedHWID LicenseBinaryBlob + MACData []byte //[16]byte +} diff --git a/protocol/nla/cssp.go b/protocol/nla/cssp.go new file mode 100644 index 0000000..e30b14a --- /dev/null +++ b/protocol/nla/cssp.go @@ -0,0 +1,98 @@ +package nla + +import ( + "encoding/asn1" + "log/slog" +) + +type NegoToken struct { + Data []byte `asn1:"explicit,tag:0"` +} + +type TSRequest struct { + Version int `asn1:"explicit,tag:0"` + NegoTokens []NegoToken `asn1:"optional,explicit,tag:1"` + AuthInfo []byte `asn1:"optional,explicit,tag:2"` + PubKeyAuth []byte `asn1:"optional,explicit,tag:3"` + //ErrorCode int `asn1:"optional,explicit,tag:4"` +} + +type TSCredentials struct { + CredType int `asn1:"explicit,tag:0"` + Credentials []byte `asn1:"explicit,tag:1"` +} + +type TSPasswordCreds struct { + DomainName []byte `asn1:"explicit,tag:0"` + UserName []byte `asn1:"explicit,tag:1"` + Password []byte `asn1:"explicit,tag:2"` +} + +type TSCspDataDetail struct { + KeySpec int `asn1:"explicit,tag:0"` + CardName string `asn1:"explicit,tag:1"` + ReaderName string `asn1:"explicit,tag:2"` + ContainerName string `asn1:"explicit,tag:3"` + CspName string `asn1:"explicit,tag:4"` +} + +type TSSmartCardCreds struct { + Pin string `asn1:"explicit,tag:0"` + CspData []TSCspDataDetail `asn1:"explicit,tag:1"` + UserHint string `asn1:"explicit,tag:2"` + DomainHint string `asn1:"explicit,tag:3"` +} + +func EncodeDERTRequest(msgs []Message, authInfo []byte, pubKeyAuth []byte) []byte { + req := TSRequest{ + Version: 2, + } + + if len(msgs) > 0 { + req.NegoTokens = make([]NegoToken, 0, len(msgs)) + } + + for _, msg := range msgs { + token := NegoToken{msg.Serialize()} + req.NegoTokens = append(req.NegoTokens, token) + } + + if len(authInfo) > 0 { + req.AuthInfo = authInfo + } + + if len(pubKeyAuth) > 0 { + req.PubKeyAuth = pubKeyAuth + } + + result, err := asn1.Marshal(req) + if err != nil { + slog.Error("EncodeDERTRequest", "err", err) + } + return result +} + +func DecodeDERTRequest(s []byte) (*TSRequest, error) { + treq := &TSRequest{} + _, err := asn1.Unmarshal(s, treq) + return treq, err +} +func EncodeDERTCredentials(domain, username, password []byte) []byte { + tpas := TSPasswordCreds{domain, username, password} + result, err := asn1.Marshal(tpas) + if err != nil { + slog.Error("EncodeDERTCredentials", "err", err) + } + tcre := TSCredentials{1, result} + result, err = asn1.Marshal(tcre) + if err != nil { + slog.Error("EncodeDERTCredentials", "err", err) + } + return result +} + +func DecodeDERTCredentials(s []byte) (*TSCredentials, error) { + tcre := &TSCredentials{} + _, err := asn1.Unmarshal(s, tcre) + return tcre, err +} diff --git a/protocol/nla/cssp_test.go b/protocol/nla/cssp_test.go new file mode 100644 index 0000000..300c9b7 --- /dev/null +++ b/protocol/nla/cssp_test.go @@ -0,0 +1,16 @@ +package nla_test + +import ( + "encoding/hex" + "testing" + + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" +) + +func TestEncodeDERTRequest(t *testing.T) { + ntlm := nla.NewNTLMv2("", "", "") + result := nla.EncodeDERTRequest([]nla.Message{ntlm.GetNegotiateMessage()}, []byte(""), []byte("")) + if hex.EncodeToString(result) != "302fa003020102a12830263024a02204204e544c4d53535000010000003582086000000000000000000000000000000000" { + t.Error("not equal") + } +} diff --git a/protocol/nla/encode.go b/protocol/nla/encode.go new file mode 100644 index 0000000..d65b639 --- /dev/null +++ b/protocol/nla/encode.go @@ -0,0 +1,46 @@ +package nla + +import ( + "crypto/hmac" + "crypto/md5" + "crypto/rc4" + "strings" + + "git.zeroonesoft.cn/golib/rdplib/core" + "golang.org/x/crypto/md4" +) + +func MD4(data []byte) []byte { + h := md4.New() + h.Write(data) + return h.Sum(nil) +} + +func MD5(data []byte) []byte { + h := md5.New() + h.Write(data) + return h.Sum(nil) +} + +func HMAC_MD5(key, data []byte) []byte { + h := hmac.New(md5.New, key) + h.Write(data) + return h.Sum(nil) +} + +// Version 2 of NTLM hash function +func NTOWFv2(password, user, domain string) []byte { + return HMAC_MD5(MD4(core.UnicodeEncode(password)), core.UnicodeEncode(strings.ToUpper(user)+domain)) +} + +// Same as NTOWFv2 +func LMOWFv2(password, user, domain string) []byte { + return NTOWFv2(password, user, domain) +} + +func RC4K(key, src []byte) []byte { + result := make([]byte, len(src)) + rc4obj, _ := rc4.NewCipher(key) + rc4obj.XORKeyStream(result, src) + return result +} diff --git a/protocol/nla/encode_test.go b/protocol/nla/encode_test.go new file mode 100644 index 0000000..bd0313d --- /dev/null +++ b/protocol/nla/encode_test.go @@ -0,0 +1,32 @@ +package nla_test + +import ( + "encoding/hex" + "testing" + + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" +) + +func TestNTOWFv2(t *testing.T) { + res := hex.EncodeToString(nla.NTOWFv2("", "", "")) + expected := "f4c1a15dd59d4da9bd595599220d971a" + if res != expected { + t.Error(res, "not equal to", expected) + } + + res = hex.EncodeToString(nla.NTOWFv2("user", "pwd", "dom")) + expected = "652feb8208b3a8a6264c9c5d5b820979" + if res != expected { + t.Error(res, "not equal to", expected) + } +} + +func TestRC4K(t *testing.T) { + key, _ := hex.DecodeString("55638e834ce774c100637f197bc0683f") + src, _ := hex.DecodeString("177d16086dd3f06fa8d594e3bad005b7") + res := hex.EncodeToString(nla.RC4K(key, src)) + expected := "f5ab375222707a492bd5a90705d96d1d" + if res != expected { + t.Error(res, "not equal to", expected) + } +} diff --git a/protocol/nla/ntlm.go b/protocol/nla/ntlm.go new file mode 100644 index 0000000..a503d71 --- /dev/null +++ b/protocol/nla/ntlm.go @@ -0,0 +1,530 @@ +package nla + +import ( + "bytes" + "crypto/md5" + "crypto/rc4" + "encoding/binary" + "encoding/hex" + "fmt" + "log/slog" + "time" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/core" +) + +const ( + WINDOWS_MINOR_VERSION_0 = 0x00 + WINDOWS_MINOR_VERSION_1 = 0x01 + WINDOWS_MINOR_VERSION_2 = 0x02 + WINDOWS_MINOR_VERSION_3 = 0x03 + + WINDOWS_MAJOR_VERSION_5 = 0x05 + WINDOWS_MAJOR_VERSION_6 = 0x06 + NTLMSSP_REVISION_W2K3 = 0x0F +) + +const ( + MsvAvEOL = 0x0000 + MsvAvNbComputerName = 0x0001 + MsvAvNbDomainName = 0x0002 + MsvAvDnsComputerName = 0x0003 + MsvAvDnsDomainName = 0x0004 + MsvAvDnsTreeName = 0x0005 + MsvAvFlags = 0x0006 + MsvAvTimestamp = 0x0007 + MsvAvSingleHost = 0x0008 + MsvAvTargetName = 0x0009 + MsvChannelBindings = 0x000A +) + +type AVPair struct { + Id uint16 `struc:"little"` + Len uint16 `struc:"little,sizeof=Value"` + Value []byte `struc:"little"` +} + +const ( + NTLMSSP_NEGOTIATE_56 = 0x80000000 + NTLMSSP_NEGOTIATE_KEY_EXCH = 0x40000000 + NTLMSSP_NEGOTIATE_128 = 0x20000000 + NTLMSSP_NEGOTIATE_VERSION = 0x02000000 + NTLMSSP_NEGOTIATE_TARGET_INFO = 0x00800000 + NTLMSSP_REQUEST_NON_NT_SESSION_KEY = 0x00400000 + NTLMSSP_NEGOTIATE_IDENTIFY = 0x00100000 + NTLMSSP_NEGOTIATE_EXTENDED_SESSIONSECURITY = 0x00080000 + NTLMSSP_TARGET_TYPE_SERVER = 0x00020000 + NTLMSSP_TARGET_TYPE_DOMAIN = 0x00010000 + NTLMSSP_NEGOTIATE_ALWAYS_SIGN = 0x00008000 + NTLMSSP_NEGOTIATE_OEM_WORKSTATION_SUPPLIED = 0x00002000 + NTLMSSP_NEGOTIATE_OEM_DOMAIN_SUPPLIED = 0x00001000 + NTLMSSP_NEGOTIATE_NTLM = 0x00000200 + NTLMSSP_NEGOTIATE_LM_KEY = 0x00000080 + NTLMSSP_NEGOTIATE_DATAGRAM = 0x00000040 + NTLMSSP_NEGOTIATE_SEAL = 0x00000020 + NTLMSSP_NEGOTIATE_SIGN = 0x00000010 + NTLMSSP_REQUEST_TARGET = 0x00000004 + NTLM_NEGOTIATE_OEM = 0x00000002 + NTLMSSP_NEGOTIATE_UNICODE = 0x00000001 +) + +type NVersion struct { + ProductMajorVersion uint8 `struc:"little"` + ProductMinorVersion uint8 `struc:"little"` + ProductBuild uint16 `struc:"little"` + Reserved [3]byte `struc:"little"` + NTLMRevisionCurrent uint8 `struc:"little"` +} + +func NewNVersion() NVersion { + return NVersion{ + ProductMajorVersion: WINDOWS_MAJOR_VERSION_6, + ProductMinorVersion: WINDOWS_MINOR_VERSION_0, + ProductBuild: 6002, + NTLMRevisionCurrent: NTLMSSP_REVISION_W2K3, + } +} + +type Message interface { + Serialize() []byte +} + +type NegotiateMessage struct { + Signature [8]byte `struc:"little"` + MessageType uint32 `struc:"little"` + NegotiateFlags uint32 `struc:"little"` + DomainNameLen uint16 `struc:"little"` + DomainNameMaxLen uint16 `struc:"little"` + DomainNameBufferOffset uint32 `struc:"little"` + WorkstationLen uint16 `struc:"little"` + WorkstationMaxLen uint16 `struc:"little"` + WorkstationBufferOffset uint32 `struc:"little"` + Version NVersion `struc:"skip"` + Payload [32]byte `struc:"skip"` +} + +func NewNegotiateMessage() *NegotiateMessage { + return &NegotiateMessage{ + Signature: [8]byte{'N', 'T', 'L', 'M', 'S', 'S', 'P', 0x00}, + MessageType: 0x00000001, + } +} + +func (m *NegotiateMessage) Serialize() []byte { + if (m.NegotiateFlags & NTLMSSP_NEGOTIATE_VERSION) != 0 { + m.Version = NewNVersion() + } + buff := &bytes.Buffer{} + struc.Pack(buff, m) + + return buff.Bytes() +} + +type ChallengeMessage struct { + totalLen int + Signature []byte `struc:"[8]byte"` + MessageType uint32 `struc:"little"` + TargetNameLen uint16 `struc:"little"` + TargetNameMaxLen uint16 `struc:"little"` + TargetNameBufferOffset uint32 `struc:"little"` + NegotiateFlags uint32 `struc:"little"` + ServerChallenge [8]byte `struc:"little"` + Reserved [8]byte `struc:"little"` + TargetInfoLen uint16 `struc:"little"` + TargetInfoMaxLen uint16 `struc:"little"` + TargetInfoBufferOffset uint32 `struc:"little"` + Version NVersion `struc:"skip"` + Payload []byte `struc:"skip"` +} + +func (m *ChallengeMessage) Serialize() []byte { + buff := &bytes.Buffer{} + struc.Pack(buff, m) + if (m.NegotiateFlags & NTLMSSP_NEGOTIATE_VERSION) != 0 { + struc.Pack(buff, m.Version) + } + buff.Write(m.Payload) + return buff.Bytes() +} + +func NewChallengeMessage() *ChallengeMessage { + return &ChallengeMessage{ + Signature: []byte{'N', 'T', 'L', 'M', 'S', 'S', 'P', 0x00}, + MessageType: 0x00000002, + } +} + +// total len - payload len +func (m *ChallengeMessage) BaseLen() uint32 { + return uint32(m.totalLen - len(m.Payload)) +} + +func (m *ChallengeMessage) getTargetInfo() []byte { + if m.TargetInfoLen == 0 { + return make([]byte, 0) + } + offset := m.BaseLen() + start := m.TargetInfoBufferOffset - offset + return m.Payload[start : start+uint32(m.TargetInfoLen)] +} +func (m *ChallengeMessage) getTargetName() []byte { + if m.TargetNameLen == 0 { + return make([]byte, 0) + } + offset := m.BaseLen() + start := m.TargetNameBufferOffset - offset + return m.Payload[start : start+uint32(m.TargetNameLen)] +} +func (m *ChallengeMessage) getTargetInfoTimestamp(data []byte) []byte { + r := bytes.NewReader(data) + for r.Len() > 0 { + avPair := &AVPair{} + struc.Unpack(r, avPair) + if avPair.Id == MsvAvTimestamp { + return avPair.Value + } + + if avPair.Id == MsvAvEOL { + break + } + } + return nil +} + +type AuthenticateMessage struct { + Signature [8]byte + MessageType uint32 `struc:"little"` + LmChallengeResponseLen uint16 `struc:"little"` + LmChallengeResponseMaxLen uint16 `struc:"little"` + LmChallengeResponseBufferOffset uint32 `struc:"little"` + NtChallengeResponseLen uint16 `struc:"little"` + NtChallengeResponseMaxLen uint16 `struc:"little"` + NtChallengeResponseBufferOffset uint32 `struc:"little"` + DomainNameLen uint16 `struc:"little"` + DomainNameMaxLen uint16 `struc:"little"` + DomainNameBufferOffset uint32 `struc:"little"` + UserNameLen uint16 `struc:"little"` + UserNameMaxLen uint16 `struc:"little"` + UserNameBufferOffset uint32 `struc:"little"` + WorkstationLen uint16 `struc:"little"` + WorkstationMaxLen uint16 `struc:"little"` + WorkstationBufferOffset uint32 `struc:"little"` + EncryptedRandomSessionLen uint16 `struc:"little"` + EncryptedRandomSessionMaxLen uint16 `struc:"little"` + EncryptedRandomSessionBufferOffset uint32 `struc:"little"` + NegotiateFlags uint32 `struc:"little"` + Version NVersion `struc:"little"` + MIC [16]byte `struc:"little"` + Payload []byte `struc:"skip"` +} + +func (m *AuthenticateMessage) BaseLen() uint32 { + return 88 +} + +func NewAuthenticateMessage(negFlag uint32, domain, user, workstation []byte, + lmchallResp, ntchallResp, enRandomSessKey []byte) *AuthenticateMessage { + msg := &AuthenticateMessage{ + Signature: [8]byte{'N', 'T', 'L', 'M', 'S', 'S', 'P', 0x00}, + MessageType: 0x00000003, + NegotiateFlags: negFlag, + } + payload := make([]byte, 0, len(lmchallResp)+len(ntchallResp)+len(domain)+len(user)+len(workstation)+len(enRandomSessKey)) + + msg.LmChallengeResponseLen = uint16(len(lmchallResp)) + msg.LmChallengeResponseMaxLen = msg.LmChallengeResponseLen + msg.LmChallengeResponseBufferOffset = msg.BaseLen() + payload = append(payload, lmchallResp...) + + msg.NtChallengeResponseLen = uint16(len(ntchallResp)) + msg.NtChallengeResponseMaxLen = msg.NtChallengeResponseLen + msg.NtChallengeResponseBufferOffset = msg.LmChallengeResponseBufferOffset + uint32(msg.LmChallengeResponseLen) + payload = append(payload, ntchallResp...) + + msg.DomainNameLen = uint16(len(domain)) + msg.DomainNameMaxLen = msg.DomainNameLen + msg.DomainNameBufferOffset = msg.NtChallengeResponseBufferOffset + uint32(msg.NtChallengeResponseLen) + payload = append(payload, domain...) + + msg.UserNameLen = uint16(len(user)) + msg.UserNameMaxLen = msg.UserNameLen + msg.UserNameBufferOffset = msg.DomainNameBufferOffset + uint32(msg.DomainNameLen) + payload = append(payload, user...) + + msg.WorkstationLen = uint16(len(workstation)) + msg.WorkstationMaxLen = msg.WorkstationLen + msg.WorkstationBufferOffset = msg.UserNameBufferOffset + uint32(msg.UserNameLen) + payload = append(payload, workstation...) + + msg.EncryptedRandomSessionLen = uint16(len(enRandomSessKey)) + msg.EncryptedRandomSessionMaxLen = msg.EncryptedRandomSessionLen + msg.EncryptedRandomSessionBufferOffset = msg.WorkstationBufferOffset + uint32(msg.WorkstationLen) + payload = append(payload, enRandomSessKey...) + + if (msg.NegotiateFlags & NTLMSSP_NEGOTIATE_VERSION) != 0 { + msg.Version = NewNVersion() + } + msg.Payload = payload + + return msg +} + +func (m *AuthenticateMessage) Serialize() []byte { + buff := &bytes.Buffer{} + struc.Pack(buff, m) + buff.Write(m.Payload) + return buff.Bytes() +} + +type NTLMv2 struct { + domain string + user string + password string + respKeyNT []byte + respKeyLM []byte + negotiateMessage *NegotiateMessage + challengeMessage *ChallengeMessage + authenticateMessage *AuthenticateMessage + enableUnicode bool +} + +func NewNTLMv2(domain, user, password string) *NTLMv2 { + return &NTLMv2{ + domain: domain, + user: user, + password: password, + respKeyNT: NTOWFv2(password, user, domain), + respKeyLM: LMOWFv2(password, user, domain), + } +} + +// generate first handshake messgae +func (n *NTLMv2) GetNegotiateMessage() *NegotiateMessage { + negoMsg := NewNegotiateMessage() + negoMsg.NegotiateFlags = NTLMSSP_NEGOTIATE_KEY_EXCH | + NTLMSSP_NEGOTIATE_128 | + NTLMSSP_NEGOTIATE_EXTENDED_SESSIONSECURITY | + NTLMSSP_NEGOTIATE_ALWAYS_SIGN | + NTLMSSP_NEGOTIATE_NTLM | + NTLMSSP_NEGOTIATE_SEAL | + NTLMSSP_NEGOTIATE_SIGN | + NTLMSSP_REQUEST_TARGET | + NTLMSSP_NEGOTIATE_UNICODE + n.negotiateMessage = negoMsg + return n.negotiateMessage +} + +// process NTLMv2 Authenticate hash +func (n *NTLMv2) ComputeResponseV2(respKeyNT, respKeyLM, serverChallenge, clientChallenge, + timestamp, serverInfo []byte) (ntChallResp, lmChallResp, SessBaseKey []byte) { + + // Build the temp blob: 2+6+8+8+4+len(serverInfo) bytes + temp := make([]byte, 0, 28+len(serverInfo)) + temp = append(temp, 0x01, 0x01) // Responser version, HiResponser version + temp = append(temp, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00) + temp = append(temp, timestamp...) + temp = append(temp, clientChallenge...) + temp = append(temp, 0x00, 0x00, 0x00, 0x00) + temp = append(temp, serverInfo...) + + ntInput := make([]byte, 0, len(serverChallenge)+len(temp)) + ntInput = append(ntInput, serverChallenge...) + ntInput = append(ntInput, temp...) + ntProof := HMAC_MD5(respKeyNT, ntInput) + + ntChallResp = make([]byte, 0, len(ntProof)+len(temp)) + ntChallResp = append(ntChallResp, ntProof...) + ntChallResp = append(ntChallResp, temp...) + + lmInput := make([]byte, 0, len(serverChallenge)+len(clientChallenge)) + lmInput = append(lmInput, serverChallenge...) + lmInput = append(lmInput, clientChallenge...) + lmChallResp = HMAC_MD5(respKeyLM, lmInput) + lmChallResp = append(lmChallResp, clientChallenge...) + + SessBaseKey = HMAC_MD5(respKeyNT, ntProof) + return +} + +func MIC(exportedSessionKey []byte, negotiateMessage, challengeMessage, authenticateMessage Message) []byte { + neg := negotiateMessage.Serialize() + chal := challengeMessage.Serialize() + auth := authenticateMessage.Serialize() + data := make([]byte, 0, len(neg)+len(chal)+len(auth)) + data = append(data, neg...) + data = append(data, chal...) + data = append(data, auth...) + return HMAC_MD5(exportedSessionKey, data) +} + +func concat(bs ...[]byte) []byte { + return bytes.Join(bs, nil) +} + +var ( + clientSigning = concat([]byte("session key to client-to-server signing key magic constant"), []byte{0x00}) + serverSigning = concat([]byte("session key to server-to-client signing key magic constant"), []byte{0x00}) + clientSealing = concat([]byte("session key to client-to-server sealing key magic constant"), []byte{0x00}) + serverSealing = concat([]byte("session key to server-to-client sealing key magic constant"), []byte{0x00}) +) + +func (n *NTLMv2) GetAuthenticateMessage(s []byte) (*AuthenticateMessage, *NTLMv2Security) { + slog.Debug("GetAuthenticateMessage", "s", s) + + challengeMsg := &ChallengeMessage{totalLen: len(s)} + r := bytes.NewReader(s) + err := struc.Unpack(r, challengeMsg) + if err != nil { + slog.Error("GetAuthenticateMessage", "err", err) + return nil, nil + } + if challengeMsg.NegotiateFlags&NTLMSSP_NEGOTIATE_VERSION != 0 { + version := NVersion{} + err := struc.Unpack(r, &version) + if err != nil { + slog.Error("GetAuthenticateMessage", "err", err) + return nil, nil + } + challengeMsg.Version = version + } + challengeMsg.Payload, _ = core.ReadBytes(r.Len(), r) + n.challengeMessage = challengeMsg + slog.Debug("GetAuthenticateMessage", "challengeMsg", challengeMsg) + + serverName := challengeMsg.getTargetName() + serverInfo := challengeMsg.getTargetInfo() + timestamp := challengeMsg.getTargetInfoTimestamp(serverInfo) + computeMIC := false + if timestamp == nil { + ft := uint64(time.Now().UnixNano()) / 100 + ft += 116444736000000000 // add time between unix & windows offset + timestamp = make([]byte, 8) + binary.LittleEndian.PutUint64(timestamp, ft) + } else { + computeMIC = true + } + slog.Debug("GetAuthenticateMessage", "serverName", core.UnicodeDecode(serverName)) + serverChallenge := challengeMsg.ServerChallenge[:] + clientChallenge := core.Random(8) + ntChallengeResponse, lmChallengeResponse, SessionBaseKey := n.ComputeResponseV2( + n.respKeyNT, n.respKeyLM, serverChallenge, clientChallenge, timestamp, serverInfo) + + exchangeKey := SessionBaseKey + exportedSessionKey := core.Random(16) + EncryptedRandomSessionKey := make([]byte, len(exportedSessionKey)) + rc, _ := rc4.NewCipher(exchangeKey) + rc.XORKeyStream(EncryptedRandomSessionKey, exportedSessionKey) + + if challengeMsg.NegotiateFlags&NTLMSSP_NEGOTIATE_UNICODE != 0 { + n.enableUnicode = true + } + slog.Debug(fmt.Sprintf("user: %s, password:********", n.user)) + domain, user, _ := n.GetEncodedCredentials() + + n.authenticateMessage = NewAuthenticateMessage(challengeMsg.NegotiateFlags, + domain, user, []byte(""), lmChallengeResponse, ntChallengeResponse, EncryptedRandomSessionKey) + + if computeMIC { + copy(n.authenticateMessage.MIC[:], MIC(exportedSessionKey, n.negotiateMessage, n.challengeMessage, n.authenticateMessage)[:16]) + } + + md := md5.New() + //ClientSigningKey + a := concat(exportedSessionKey, clientSigning) + md.Write(a) + ClientSigningKey := md.Sum(nil) + //ServerSigningKey + md.Reset() + a = concat(exportedSessionKey, serverSigning) + md.Write(a) + ServerSigningKey := md.Sum(nil) + //ClientSealingKey + md.Reset() + a = concat(exportedSessionKey, clientSealing) + md.Write(a) + ClientSealingKey := md.Sum(nil) + //ServerSealingKey + md.Reset() + a = concat(exportedSessionKey, serverSealing) + md.Write(a) + ServerSealingKey := md.Sum(nil) + + slog.Debug(fmt.Sprintf("ClientSigningKey:%s", hex.EncodeToString(ClientSigningKey))) + slog.Debug(fmt.Sprintf("ServerSigningKey:%s", hex.EncodeToString(ServerSigningKey))) + slog.Debug(fmt.Sprintf("ClientSealingKey:%s", hex.EncodeToString(ClientSealingKey))) + slog.Debug(fmt.Sprintf("ServerSealingKey:%s", hex.EncodeToString(ServerSealingKey))) + + encryptRC4, _ := rc4.NewCipher(ClientSealingKey) + decryptRC4, _ := rc4.NewCipher(ServerSealingKey) + + ntlmSec := &NTLMv2Security{encryptRC4, decryptRC4, ClientSigningKey, ServerSigningKey, 0} + + return n.authenticateMessage, ntlmSec +} + +func (n *NTLMv2) GetEncodedCredentials() ([]byte, []byte, []byte) { + if n.enableUnicode { + return core.UnicodeEncode(n.domain), core.UnicodeEncode(n.user), core.UnicodeEncode(n.password) + } + return []byte(n.domain), []byte(n.user), []byte(n.password) +} + +type NTLMv2Security struct { + EncryptRC4 *rc4.Cipher + DecryptRC4 *rc4.Cipher + SigningKey []byte + VerifyKey []byte + SeqNum uint32 +} + +func (n *NTLMv2Security) GssEncrypt(s []byte) []byte { + p := make([]byte, len(s)) + n.EncryptRC4.XORKeyStream(p, s) + + // HMAC input: SeqNum(4) + plaintext + sigInput := make([]byte, 4+len(s)) + binary.LittleEndian.PutUint32(sigInput, n.SeqNum) + copy(sigInput[4:], s) + s1 := HMAC_MD5(n.SigningKey, sigInput)[:8] + + checksum := make([]byte, 8) + n.EncryptRC4.XORKeyStream(checksum, s1) + + // Output: version(4) + checksum(8) + SeqNum(4) + encrypted(len(p)) + out := make([]byte, 16+len(p)) + binary.LittleEndian.PutUint32(out[0:], 0x00000001) + copy(out[4:], checksum) + binary.LittleEndian.PutUint32(out[12:], n.SeqNum) + copy(out[16:], p) + + n.SeqNum++ + return out +} + +func (n *NTLMv2Security) GssDecrypt(s []byte) []byte { + if len(s) < 16 { + return nil + } + // s[0:4] = version (ignored), s[4:12] = checksum, s[12:16] = seqNum, s[16:] = data + checksum := s[4:12] + seqNum := binary.LittleEndian.Uint32(s[12:16]) + data := s[16:] + + p := make([]byte, len(data)) + n.DecryptRC4.XORKeyStream(p, data) + + check := make([]byte, 8) + n.DecryptRC4.XORKeyStream(check, checksum) + + // HMAC input: seqNum(4) + decrypted(len(p)) + verifyInput := make([]byte, 4+len(p)) + binary.LittleEndian.PutUint32(verifyInput, seqNum) + copy(verifyInput[4:], p) + verify := HMAC_MD5(n.VerifyKey, verifyInput)[:8] + + if !bytes.Equal(verify, check) { + return nil + } + return p +} diff --git a/protocol/nla/ntlm_test.go b/protocol/nla/ntlm_test.go new file mode 100644 index 0000000..be68d76 --- /dev/null +++ b/protocol/nla/ntlm_test.go @@ -0,0 +1,53 @@ +package nla_test + +import ( + "bytes" + "encoding/hex" + "testing" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" +) + +func TestNewNegotiateMessage(t *testing.T) { + ntlm := nla.NewNTLMv2("", "", "") + negoMsg := ntlm.GetNegotiateMessage() + buff := &bytes.Buffer{} + struc.Pack(buff, negoMsg) + + result := hex.EncodeToString(buff.Bytes()) + expected := "4e544c4d535350000100000035820860000000000000000000000000000000000000000000000000" + + if result != expected { + t.Error(result, " not equals to", expected) + } +} + +func TestNTLMv2_ComputeResponse(t *testing.T) { + ntlm := nla.NewNTLMv2("", "", "") + + ResponseKeyNT, _ := hex.DecodeString("39e32c766260586a9036f1ceb04c3007") + ResponseKeyLM, _ := hex.DecodeString("39e32c766260586a9036f1ceb04c3007") + ServerChallenge, _ := hex.DecodeString("adcb9d1c8d4a5ed8") + ClienChallenge, _ := hex.DecodeString("1a78bed8e5d5efa7") + Timestamp, _ := hex.DecodeString("a02f44f01267d501") + ServerName, _ := hex.DecodeString("02001e00570049004e002d00460037005200410041004d004100500034004a00430001001e00570049004e002d00460037005200410041004d004100500034004a00430004001e00570049004e002d00460037005200410041004d004100500034004a00430003001e00570049004e002d00460037005200410041004d004100500034004a00430007000800a02f44f01267d50100000000") + + NtChallengeResponse, LmChallengeResponse, SessionBaseKey := ntlm.ComputeResponseV2(ResponseKeyNT, ResponseKeyLM, ServerChallenge, ClienChallenge, Timestamp, ServerName) + + ntChallRespExpected := "4e7316531937d2fc91e7230853844b890101000000000000a02f44f01267d5011a78bed8e5d5efa70000000002001e00570049004e002d00460037005200410041004d004100500034004a00430001001e00570049004e002d00460037005200410041004d004100500034004a00430004001e00570049004e002d00460037005200410041004d004100500034004a00430003001e00570049004e002d00460037005200410041004d004100500034004a00430007000800a02f44f01267d50100000000" + lmChallRespExpected := "d4dc6edc0c37dd70f69b5c4f05a615661a78bed8e5d5efa7" + sessBaseKeyExpected := "034009be89a0507b2bd6d28e966e1dab" + + if hex.EncodeToString(NtChallengeResponse) != ntChallRespExpected { + t.Error("NtChallengeResponse incorrect") + } + + if hex.EncodeToString(LmChallengeResponse) != lmChallRespExpected { + t.Error("LmChallengeResponse incorrect") + } + + if hex.EncodeToString(SessionBaseKey) != sessBaseKeyExpected { + t.Error("SessionBaseKey incorrect") + } +} diff --git a/protocol/pdu/caps.go b/protocol/pdu/caps.go new file mode 100644 index 0000000..1802273 --- /dev/null +++ b/protocol/pdu/caps.go @@ -0,0 +1,857 @@ +package pdu + +import ( + "bytes" + "encoding/hex" + "errors" + "fmt" + "io" + "log/slog" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/gcc" +) + +type CapsType uint16 + +const ( + CAPSTYPE_GENERAL CapsType = 0x0001 + CAPSTYPE_BITMAP = 0x0002 + CAPSTYPE_ORDER = 0x0003 + CAPSTYPE_BITMAPCACHE = 0x0004 + CAPSTYPE_CONTROL = 0x0005 + CAPSTYPE_ACTIVATION = 0x0007 + CAPSTYPE_POINTER = 0x0008 + CAPSTYPE_SHARE = 0x0009 + CAPSTYPE_COLORCACHE = 0x000A + CAPSTYPE_SOUND = 0x000C + CAPSTYPE_INPUT = 0x000D + CAPSTYPE_FONT = 0x000E + CAPSTYPE_BRUSH = 0x000F + CAPSTYPE_GLYPHCACHE = 0x0010 + CAPSTYPE_OFFSCREENCACHE = 0x0011 + CAPSTYPE_BITMAPCACHE_HOSTSUPPORT = 0x0012 + CAPSTYPE_BITMAPCACHE_REV2 = 0x0013 + CAPSTYPE_VIRTUALCHANNEL = 0x0014 + CAPSTYPE_DRAWNINEGRIDCACHE = 0x0015 + CAPSTYPE_DRAWGDIPLUS = 0x0016 + CAPSTYPE_RAIL = 0x0017 + CAPSTYPE_WINDOW = 0x0018 + CAPSETTYPE_COMPDESK = 0x0019 + CAPSETTYPE_MULTIFRAGMENTUPDATE = 0x001A + CAPSETTYPE_LARGE_POINTER = 0x001B + CAPSETTYPE_SURFACE_COMMANDS = 0x001C + CAPSETTYPE_BITMAP_CODECS = 0x001D + CAPSSETTYPE_FRAME_ACKNOWLEDGE = 0x001E +) + +func (c CapsType) String() string { + switch c { + case CAPSTYPE_GENERAL: + return "CAPSTYPE_GENERAL" + case CAPSTYPE_BITMAP: + return "CAPSTYPE_BITMAP" + case CAPSTYPE_ORDER: + return "CAPSTYPE_ORDER" + case CAPSTYPE_BITMAPCACHE: + return "CAPSTYPE_BITMAPCACHE" + case CAPSTYPE_CONTROL: + return "CAPSTYPE_CONTROL" + case CAPSTYPE_ACTIVATION: + return "CAPSTYPE_ACTIVATION" + case CAPSTYPE_POINTER: + return "CAPSTYPE_POINTER" + case CAPSTYPE_SHARE: + return "CAPSTYPE_SHARE" + case CAPSTYPE_COLORCACHE: + return "CAPSTYPE_COLORCACHE" + case CAPSTYPE_SOUND: + return "CAPSTYPE_SOUND" + case CAPSTYPE_INPUT: + return "CAPSTYPE_INPUT" + case CAPSTYPE_FONT: + return "CAPSTYPE_FONT" + case CAPSTYPE_BRUSH: + return "CAPSTYPE_BRUSH" + case CAPSTYPE_GLYPHCACHE: + return "CAPSTYPE_GLYPHCACHE" + case CAPSTYPE_OFFSCREENCACHE: + return "CAPSTYPE_OFFSCREENCACHE" + case CAPSTYPE_BITMAPCACHE_HOSTSUPPORT: + return "CAPSTYPE_BITMAPCACHE_HOSTSUPPORT" + case CAPSTYPE_BITMAPCACHE_REV2: + return "CAPSTYPE_BITMAPCACHE_REV2" + case CAPSTYPE_VIRTUALCHANNEL: + return "CAPSTYPE_VIRTUALCHANNEL" + case CAPSTYPE_DRAWNINEGRIDCACHE: + return "CAPSTYPE_DRAWNINEGRIDCACHE" + case CAPSTYPE_DRAWGDIPLUS: + return "CAPSTYPE_DRAWGDIPLUS" + case CAPSTYPE_RAIL: + return "CAPSTYPE_RAIL" + case CAPSTYPE_WINDOW: + return "CAPSTYPE_WINDOW" + case CAPSETTYPE_COMPDESK: + return "CAPSETTYPE_COMPDESK" + case CAPSETTYPE_MULTIFRAGMENTUPDATE: + return "CAPSETTYPE_MULTIFRAGMENTUPDATE" + case CAPSETTYPE_LARGE_POINTER: + return "CAPSETTYPE_LARGE_POINTER" + case CAPSETTYPE_SURFACE_COMMANDS: + return "CAPSETTYPE_SURFACE_COMMANDS" + case CAPSETTYPE_BITMAP_CODECS: + return "CAPSETTYPE_BITMAP_CODECS" + case CAPSSETTYPE_FRAME_ACKNOWLEDGE: + return "CAPSSETTYPE_FRAME_ACKNOWLEDGE" + } + + return "Unknown" +} + +type MajorType uint16 + +const ( + OSMAJORTYPE_UNSPECIFIED MajorType = 0x0000 + OSMAJORTYPE_WINDOWS = 0x0001 + OSMAJORTYPE_OS2 = 0x0002 + OSMAJORTYPE_MACINTOSH = 0x0003 + OSMAJORTYPE_UNIX = 0x0004 + OSMAJORTYPE_IOS = 0x0005 + OSMAJORTYPE_OSX = 0x0006 + OSMAJORTYPE_ANDROID = 0x0007 +) + +type MinorType uint16 + +const ( + OSMINORTYPE_UNSPECIFIED MinorType = 0x0000 + OSMINORTYPE_WINDOWS_31X = 0x0001 + OSMINORTYPE_WINDOWS_95 = 0x0002 + OSMINORTYPE_WINDOWS_NT = 0x0003 + OSMINORTYPE_OS2_V21 = 0x0004 + OSMINORTYPE_POWER_PC = 0x0005 + OSMINORTYPE_MACINTOSH = 0x0006 + OSMINORTYPE_NATIVE_XSERVER = 0x0007 + OSMINORTYPE_PSEUDO_XSERVER = 0x0008 + OSMINORTYPE_WINDOWS_RT = 0x0009 +) + +const ( + FASTPATH_OUTPUT_SUPPORTED uint16 = 0x0001 + NO_BITMAP_COMPRESSION_HDR = 0x0400 + LONG_CREDENTIALS_SUPPORTED = 0x0004 + AUTORECONNECT_SUPPORTED = 0x0008 + ENC_SALTED_CHECKSUM = 0x0010 +) + +type OrderFlag uint16 + +const ( + NEGOTIATEORDERSUPPORT OrderFlag = 0x0002 + ZEROBOUNDSDELTASSUPPORT = 0x0008 + COLORINDEXSUPPORT = 0x0020 + SOLIDPATTERNBRUSHONLY = 0x0040 + ORDERFLAGS_EXTRA_FLAGS = 0x0080 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240556.aspx + */ +type Order uint8 + +const ( + TS_NEG_DSTBLT_INDEX Order = 0x00 + TS_NEG_PATBLT_INDEX = 0x01 + TS_NEG_SCRBLT_INDEX = 0x02 + TS_NEG_MEMBLT_INDEX = 0x03 + TS_NEG_MEM3BLT_INDEX = 0x04 + TS_NEG_DRAWNINEGRID_INDEX = 0x07 + TS_NEG_LINETO_INDEX = 0x08 + TS_NEG_MULTI_DRAWNINEGRID_INDEX = 0x09 + TS_NEG_SAVEBITMAP_INDEX = 0x0B + TS_NEG_MULTIDSTBLT_INDEX = 0x0F + TS_NEG_MULTIPATBLT_INDEX = 0x10 + TS_NEG_MULTISCRBLT_INDEX = 0x11 + TS_NEG_MULTIOPAQUERECT_INDEX = 0x12 + TS_NEG_FAST_INDEX_INDEX = 0x13 + TS_NEG_POLYGON_SC_INDEX = 0x14 + TS_NEG_POLYGON_CB_INDEX = 0x15 + TS_NEG_POLYLINE_INDEX = 0x16 + TS_NEG_FAST_GLYPH_INDEX = 0x18 + TS_NEG_ELLIPSE_SC_INDEX = 0x19 + TS_NEG_ELLIPSE_CB_INDEX = 0x1A + TS_NEG_GLYPH_INDEX_INDEX = 0x1B +) + +type OrderEx uint16 + +const ( + ORDERFLAGS_EX_CACHE_BITMAP_REV3_SUPPORT OrderEx = 0x0002 + ORDERFLAGS_EX_ALTSEC_FRAME_MARKER_SUPPORT = 0x0004 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240563.aspx + */ + +const ( + INPUT_FLAG_SCANCODES uint16 = 0x0001 + INPUT_FLAG_MOUSEX = 0x0004 + INPUT_FLAG_FASTPATH_INPUT = 0x0008 + INPUT_FLAG_UNICODE = 0x0010 + INPUT_FLAG_FASTPATH_INPUT2 = 0x0020 + INPUT_FLAG_UNUSED1 = 0x0040 + INPUT_FLAG_UNUSED2 = 0x0080 + INPUT_FLAG_MOUSE_HWHEEL = 0x0100 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240564.aspx + */ +type BrushSupport uint32 + +const ( + BRUSH_DEFAULT BrushSupport = 0x00000000 + BRUSH_COLOR_8x8 = 0x00000001 + BRUSH_COLOR_FULL = 0x00000002 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240565.aspx + */ +type GlyphSupport uint16 + +const ( + GLYPH_SUPPORT_NONE GlyphSupport = 0x0000 + GLYPH_SUPPORT_PARTIAL = 0x0001 + GLYPH_SUPPORT_FULL = 0x0002 + GLYPH_SUPPORT_ENCODE = 0x0003 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240550.aspx + */ +type OffscreenSupportLevel uint32 + +const ( + OSL_FALSE OffscreenSupportLevel = 0x00000000 + OSL_TRUE = 0x00000001 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240551.aspx + */ +type VirtualChannelCompressionFlag uint32 + +const ( + VCCAPS_NO_COMPR VirtualChannelCompressionFlag = 0x00000000 + VCCAPS_COMPR_SC = 0x00000001 + VCCAPS_COMPR_CS_8K = 0x00000002 +) + +type SoundFlag uint16 + +const ( + SOUND_NONE SoundFlag = 0x0000 + SOUND_BEEPS_FLAG = 0x0001 +) + +type RailsupportLevel uint32 + +const ( + RAIL_LEVEL_SUPPORTED = 0x00000001 + RAIL_LEVEL_DOCKED_LANGBAR_SUPPORTED = 0x00000002 + RAIL_LEVEL_SHELL_INTEGRATION_SUPPORTED = 0x00000004 + RAIL_LEVEL_LANGUAGE_IME_SYNC_SUPPORTED = 0x00000008 + RAIL_LEVEL_SERVER_TO_CLIENT_IME_SYNC_SUPPORTED = 0x00000010 + RAIL_LEVEL_HIDE_MINIMIZED_APPS_SUPPORTED = 0x00000020 + RAIL_LEVEL_WINDOW_CLOAKING_SUPPORTED = 0x00000040 + RAIL_LEVEL_HANDSHAKE_EX_SUPPORTED = 0x00000080 +) + +const ( + INPUT_EVENT_SYNC = 0x0000 + INPUT_EVENT_UNUSED = 0x0002 + INPUT_EVENT_SCANCODE = 0x0004 + INPUT_EVENT_UNICODE = 0x0005 + INPUT_EVENT_MOUSE = 0x8001 + INPUT_EVENT_MOUSEX = 0x8002 +) + +const ( + PTRFLAGS_HWHEEL = 0x0400 + PTRFLAGS_WHEEL = 0x0200 + PTRFLAGS_WHEEL_NEGATIVE = 0x0100 + WheelRotationMask = 0x01FF + PTRFLAGS_MOVE = 0x0800 + PTRFLAGS_DOWN = 0x8000 + PTRFLAGS_BUTTON1 = 0x1000 + PTRFLAGS_BUTTON2 = 0x2000 + PTRFLAGS_BUTTON3 = 0x4000 +) + +const ( + KBDFLAGS_EXTENDED = 0x0100 + KBDFLAGS_EXTENDED1 = 0x0200 + KBDFLAGS_DOWN = 0x4000 + KBDFLAGS_RELEASE = 0x8000 +) + +// Fast-Path Input event codes (MS-RDPBCGR §2.2.8.1.2.2). Encoded in the +// upper 3 bits of the eventHeader byte. +const ( + FASTPATH_INPUT_EVENT_SCANCODE = 0 + FASTPATH_INPUT_EVENT_MOUSE = 1 + FASTPATH_INPUT_EVENT_MOUSEX = 2 + FASTPATH_INPUT_EVENT_SYNC = 3 + FASTPATH_INPUT_EVENT_UNICODE = 4 +) + +// Fast-Path keyboard event flags (eventHeader bits 0-4). +const ( + FASTPATH_INPUT_KBDFLAGS_RELEASE = 0x01 + FASTPATH_INPUT_KBDFLAGS_EXTENDED = 0x02 + FASTPATH_INPUT_KBDFLAGS_EXTENDED1 = 0x04 +) + +type SurfaceCmdFlags uint32 + +const ( + SURFCMDS_SET_SURFACE_BITS = 0x00000002 + SURFCMDS_FRAME_MARKER = 0x00000010 + SURFCMDS_STREAM_SURFACE_BITS = 0x00000040 +) + +type Capability interface { + Type() CapsType +} + +type GeneralCapability struct { + // 010018000100030000020000000015040000000000000000 + OSMajorType MajorType `struc:"little"` + OSMinorType MinorType `struc:"little"` + ProtocolVersion uint16 `struc:"little"` + Pad2octetsA uint16 `struc:"little"` + GeneralCompressionTypes uint16 `struc:"little"` + ExtraFlags uint16 `struc:"little"` + UpdateCapabilityFlag uint16 `struc:"little"` + RemoteUnshareFlag uint16 `struc:"little"` + GeneralCompressionLevel uint16 `struc:"little"` + RefreshRectSupport uint8 `struc:"little"` + SuppressOutputSupport uint8 `struc:"little"` +} + +func (*GeneralCapability) Type() CapsType { + return CAPSTYPE_GENERAL +} + +type BitmapCapability struct { + // 02001c00180001000100010000052003000000000100000001000000 + PreferredBitsPerPixel gcc.HighColor `struc:"little"` + Receive1BitPerPixel uint16 `struc:"little"` + Receive4BitsPerPixel uint16 `struc:"little"` + Receive8BitsPerPixel uint16 `struc:"little"` + DesktopWidth uint16 `struc:"little"` + DesktopHeight uint16 `struc:"little"` + Pad2octets uint16 `struc:"little"` + DesktopResizeFlag uint16 `struc:"little"` + BitmapCompressionFlag uint16 `struc:"little"` + HighColorFlags uint8 `struc:"little"` + DrawingFlags uint8 `struc:"little"` + MultipleRectangleSupport uint16 `struc:"little"` + Pad2octetsB uint16 `struc:"little"` +} + +func (*BitmapCapability) Type() CapsType { + return CAPSTYPE_BITMAP +} + +type BitmapCacheCapability struct { + // 04002800000000000000000000000000000000000000000000000000000000000000000000000000 + Pad1 uint32 `struc:"little"` + Pad2 uint32 `struc:"little"` + Pad3 uint32 `struc:"little"` + Pad4 uint32 `struc:"little"` + Pad5 uint32 `struc:"little"` + Pad6 uint32 `struc:"little"` + Cache0Entries uint16 `struc:"little"` + Cache0MaximumCellSize uint16 `struc:"little"` + Cache1Entries uint16 `struc:"little"` + Cache1MaximumCellSize uint16 `struc:"little"` + Cache2Entries uint16 `struc:"little"` + Cache2MaximumCellSize uint16 `struc:"little"` +} + +func (*BitmapCacheCapability) Type() CapsType { + return CAPSTYPE_BITMAPCACHE +} + +// BitmapCacheV2CellInfo 一个 v2 缓存单元的条目数与最大尺寸(KB)。 +type BitmapCacheV2CellInfo struct { + NumEntries uint16 `struc:"little"` + MaxCellSize uint16 `struc:"little"` +} + +// BitmapCacheRev2Capability 是 TS_BITMAPCACHE_REV2_CAPABILITYSET +// (MS-RDPBCGR 2.2.7.1.8,CAPSTYPE_BITMAPCACHE_REV2)。客户端声明 5 个 +// v2 缓存单元后,服务器即可用 CacheBitmapV2 次级订单存位图、用 MemBlt +// 主订单引用回贴(stage6 6.4b 会话内位图缓存)。 +type BitmapCacheRev2Capability struct { + CacheFlags [3]uint8 `struc:"little"` + Pad1 uint8 `struc:"little"` + CacheCells [5]BitmapCacheV2CellInfo `struc:"little"` +} + +func (*BitmapCacheRev2Capability) Type() CapsType { + return CAPSTYPE_BITMAPCACHE_REV2 +} + +type OrderCapability struct { + // 030058000000000000000000000000000000000000000000010014000000010000000a0000000000000000000000000000000000000000000000000000000000000000000000000000000000008403000000000000000000 + TerminalDescriptor [16]byte + Pad4octetsA uint32 `struc:"little"` + DesktopSaveXGranularity uint16 `struc:"little"` + DesktopSaveYGranularity uint16 `struc:"little"` + Pad2octetsA uint16 `struc:"little"` + MaximumOrderLevel uint16 `struc:"little"` + NumberFonts uint16 `struc:"little"` + OrderFlags OrderFlag `struc:"little"` + OrderSupport [32]byte + TextFlags uint16 `struc:"little"` + OrderSupportExFlags uint16 `struc:"little"` + Pad4octetsB uint32 `struc:"little"` + DesktopSaveSize uint32 `struc:"little"` + Pad2octetsC uint16 `struc:"little"` + Pad2octetsD uint16 `struc:"little"` + TextANSICodePage uint16 `struc:"little"` + Pad2octetsE uint16 `struc:"little"` +} + +func (*OrderCapability) Type() CapsType { + return CAPSTYPE_ORDER +} + +type PointerCapability struct { + ColorPointerFlag uint16 `struc:"little"` + ColorPointerCacheSize uint16 `struc:"little"` + // old version of rdp doesn't support ... + PointerCacheSize uint16 `struc:"little"` // only server need +} + +func (*PointerCapability) Type() CapsType { + return CAPSTYPE_POINTER +} + +type InputCapability struct { + // 0d005c001500000009040000040000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000c000000 + Flags uint16 `struc:"little"` + Pad2octetsA uint16 `struc:"little"` + // same value as gcc.ClientCoreSettings.kbdLayout + KeyboardLayout gcc.KeyboardLayout `struc:"little"` + // same value as gcc.ClientCoreSettings.keyboardType + KeyboardType uint32 `struc:"little"` + // same value as gcc.ClientCoreSettings.keyboardSubType + KeyboardSubType uint32 `struc:"little"` + // same value as gcc.ClientCoreSettings.keyboardFnKeys + KeyboardFunctionKey uint32 `struc:"little"` + // same value as gcc.ClientCoreSettingrrs.imeFileName + ImeFileName [64]byte + //need add 0c000000 in the end +} + +func (*InputCapability) Type() CapsType { + return CAPSTYPE_INPUT +} + +type BrushCapability struct { + // 0f00080000000000 + SupportLevel BrushSupport `struc:"little"` +} + +func (*BrushCapability) Type() CapsType { + return CAPSTYPE_BRUSH +} + +type cacheEntry struct { + Entries uint16 `struc:"little"` + MaximumCellSize uint16 `struc:"little"` +} + +type GlyphCapability struct { + // 10003400000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000 + GlyphCache [10]cacheEntry `struc:"little"` + FragCache uint32 `struc:"little"` + SupportLevel GlyphSupport `struc:"little"` + Pad2octets uint16 `struc:"little"` +} + +func (*GlyphCapability) Type() CapsType { + return CAPSTYPE_GLYPHCACHE +} + +type OffscreenBitmapCacheCapability struct { + // 11000c000000000000000000 + SupportLevel OffscreenSupportLevel `struc:"little"` + CacheSize uint16 `struc:"little"` + CacheEntries uint16 `struc:"little"` +} + +func (*OffscreenBitmapCacheCapability) Type() CapsType { + return CAPSTYPE_OFFSCREENCACHE +} + +type BitmapCache2Capability struct { + BitmapCachePersist uint16 `struc:"little"` + Pad2octets uint8 `struc:"little"` + CachesNum uint8 `struc:"little"` + BmpC0Cells uint32 `struc:"little"` + BmpC1Cells uint32 `struc:"little"` + BmpC2Cells uint32 `struc:"little"` + BmpC3Cells uint32 `struc:"little"` + BmpC4Cells uint32 `struc:"little"` + Pad2octets1 [12]byte `struc:"little"` +} + +func (*BitmapCache2Capability) Type() CapsType { + return CAPSTYPE_BITMAPCACHE_REV2 +} + +type VirtualChannelCapability struct { + // 14000c000000000000000000 + Flags VirtualChannelCompressionFlag `struc:"little"` + VCChunkSize uint32 `struc:"little"` // optional +} + +func (*VirtualChannelCapability) Type() CapsType { + return CAPSTYPE_VIRTUALCHANNEL +} + +type SoundCapability struct { + // 0c00080000000000 + Flags SoundFlag `struc:"little"` + Pad2octets uint16 `struc:"little"` +} + +func (*SoundCapability) Type() CapsType { + return CAPSTYPE_SOUND +} + +type ControlCapability struct { + ControlFlags uint16 `struc:"little"` + RemoteDetachFlag uint16 `struc:"little"` + ControlInterest uint16 `struc:"little"` + DetachInterest uint16 `struc:"little"` +} + +func (*ControlCapability) Type() CapsType { + return CAPSTYPE_CONTROL +} + +type WindowActivationCapability struct { + HelpKeyFlag uint16 `struc:"little"` + HelpKeyIndexFlag uint16 `struc:"little"` + HelpExtendedKeyFlag uint16 `struc:"little"` + WindowManagerKeyFlag uint16 `struc:"little"` +} + +func (*WindowActivationCapability) Type() CapsType { + return CAPSTYPE_ACTIVATION +} + +type FontCapability struct { + SupportFlags uint16 `struc:"little"` + Pad2octets uint16 `struc:"little"` +} + +func (*FontCapability) Type() CapsType { + return CAPSTYPE_FONT +} + +type ColorCacheCapability struct { + CacheSize uint16 `struc:"little"` + Pad2octets uint16 `struc:"little"` +} + +func (*ColorCacheCapability) Type() CapsType { + return CAPSTYPE_COLORCACHE +} + +type ShareCapability struct { + NodeId uint16 `struc:"little"` + Pad2octets uint16 `struc:"little"` +} + +func (*ShareCapability) Type() CapsType { + return CAPSTYPE_SHARE +} + +type MultiFragmentUpdate struct { + // 1a00080000000000 + MaxRequestSize uint32 `struc:"little"` +} + +func (*MultiFragmentUpdate) Type() CapsType { + return CAPSETTYPE_MULTIFRAGMENTUPDATE +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpegdi/52635737-d144-4f47-9c88-b48ceaf3efb4 + +type DrawGDIPlusCapability struct { + SupportLevel uint32 + GdipVersion uint32 + CacheLevel uint32 + GdipCacheEntries [10]byte + GdipCacheChunkSize [8]byte + GdipImageCacheProperties [6]byte +} + +func (*DrawGDIPlusCapability) Type() CapsType { + return CAPSTYPE_DRAWGDIPLUS +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/86507fed-a0ee-4242-b802-237534a8f65e +type BitmapCodec struct { + GUID [16]byte + ID uint8 + PropertiesLength uint16 `struc:"little,sizeof=Properties"` + Properties []byte +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/408b1878-9f6e-4106-8329-1af42219ba6a +type BitmapCodecS struct { + Count uint8 `struc:"sizeof=Array"` + Array []BitmapCodec +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/17e80f50-d163-49de-a23b-fd6456aa472f +type BitmapCodecsCapability struct { + SupportedBitmapCodecs BitmapCodecS // A variable-length field containing a TS_BITMAPCODECS structure (section 2.2.7.2.10.1). +} + +func (*BitmapCodecsCapability) Type() CapsType { + return CAPSETTYPE_BITMAP_CODECS +} + +// GUIDs for well-known bitmap codecs (wire byte order: LE for first three groups) +var ( + guidNSCodec = [16]byte{0xB9, 0x1B, 0x8D, 0xCA, 0x0F, 0x00, 0x4F, 0x15, 0x58, 0x9F, 0xAE, 0x2D, 0x1A, 0x87, 0xE2, 0xD6} + guidRemoteFX = [16]byte{0x12, 0x2F, 0x77, 0x76, 0x72, 0xBD, 0x63, 0x44, 0xAF, 0xB3, 0xB7, 0x3C, 0x9C, 0x6F, 0x78, 0x86} +) + +// rfxPropertiesBytes is the pre-computed TS_RFX_CLNT_CAPS_CONTAINER blob +// (MS-RDPRFX 2.2.1.1). The content is constant across all connections. +var rfxPropertiesBytes = []byte{ + 0x29, 0x00, 0x00, 0x00, // length = 41 + 0x00, 0x00, 0x00, 0x00, // captureFlags + 0x1D, 0x00, 0x00, 0x00, // capsLength = 29 + 0xC0, 0xCB, // blockType = CBY_CAPS + 0x08, 0x00, 0x00, 0x00, // blockLen + 0x01, 0x00, // numCapsets + 0xC1, 0xCB, // blockType = CBY_CAPSET + 0x15, 0x00, 0x00, 0x00, // blockLen = 21 + 0x01, // codecId + 0xC0, 0xCF, // capsetType = CLY_CAPSET + 0x01, 0x00, // numIcaps + 0x08, 0x00, // icapLen + 0x00, 0x01, // version + 0x40, 0x00, // tileSize = 64 + 0x01, // flags = VIDEOMODE + 0x01, // colConvBits + 0x01, // transformBits + 0x04, // entropyBits = RLGR3 +} + +// buildRemoteFxProperties constructs the TS_RFX_CLNT_CAPS_CONTAINER (MS-RDPRFX 2.2.1.1). +func buildRemoteFxProperties() []byte { + return bytes.Clone(rfxPropertiesBytes) +} + +// newClientBitmapCodecsCapability returns a BitmapCodecsCapability with +// NSCodec (ID=1) and RemoteFX (ID=3) entries, matching what FreeRDP advertises. +func newClientBitmapCodecsCapability() *BitmapCodecsCapability { + return &BitmapCodecsCapability{ + SupportedBitmapCodecs: BitmapCodecS{ + Array: []BitmapCodec{ + { + GUID: guidNSCodec, + ID: 1, + Properties: []byte{1, 1, 3}, // fAllowDynamicFidelity, fAllowSubsampling, colorLossLevel + }, + { + GUID: guidRemoteFX, + ID: 3, + Properties: buildRemoteFxProperties(), + }, + }, + }, + } +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/fc05c385-46c3-42cb-9ed2-c475a3990e0b +type BitmapCacheHostSupportCapability struct { + CacheVersion uint8 + Pad1 uint8 + Pad2 uint16 +} + +func (*BitmapCacheHostSupportCapability) Type() CapsType { + return CAPSTYPE_BITMAPCACHE_HOSTSUPPORT +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/41323437-c753-460e-8108-495a6fdd68a8 +type LargePointerCapability struct { + SupportFlags uint16 `struc:"little"` +} + +func (*LargePointerCapability) Type() CapsType { + return CAPSETTYPE_LARGE_POINTER +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdperp/36a25e21-25e1-4954-aae8-09aaf6715c79 +type RemoteProgramsCapability struct { + RailSupportLevel uint32 `struc:"little"` +} + +func (*RemoteProgramsCapability) Type() CapsType { + return CAPSTYPE_RAIL +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdperp/82ec7a69-f7e3-4294-830d-666178b35d15 +type WindowListCapability struct { + WndSupportLevel uint32 `struc:"little"` + NumIconCaches uint8 + NumIconCacheEntries uint16 `struc:"little"` +} + +func (*WindowListCapability) Type() CapsType { + return CAPSTYPE_WINDOW +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/9132002f-f133-4a0f-ba2f-2dc48f1e7f93 +type DesktopCompositionCapability struct { + CompDeskSupportLevel uint16 `struc:"little"` +} + +func (*DesktopCompositionCapability) Type() CapsType { + return CAPSETTYPE_COMPDESK +} + +// see https://docs.microsoft.com/en-us/openspecs/windows_protocols/ms-rdpbcgr/aa953018-c0a8-4761-bb12-86586c2cd56a +type SurfaceCommandsCapability struct { + CmdFlags uint32 `struc:"little"` + Reserved uint32 `struc:"little"` +} + +func (*SurfaceCommandsCapability) Type() CapsType { + return CAPSETTYPE_SURFACE_COMMANDS +} + +type FrameAcknowledgeCapability struct { + FrameCount uint32 `struc:"little"` +} + +func (*FrameAcknowledgeCapability) Type() CapsType { + return CAPSSETTYPE_FRAME_ACKNOWLEDGE +} + +type DrawNineGridCapability struct { + SupportLevel uint32 `struc:"little"` + CacheSize uint16 `struc:"little"` + CacheEntries uint16 `struc:"little"` +} + +func (*DrawNineGridCapability) Type() CapsType { + return CAPSTYPE_DRAWNINEGRIDCACHE +} + +func readCapability(r io.Reader) (Capability, error) { + capType, err := core.ReadUint16LE(r) + if err != nil { + return nil, err + } + capLen, err := core.ReadUint16LE(r) + if err != nil { + return nil, err + } + if int(capLen)-4 <= 0 { + return nil, errors.New(fmt.Sprintf("Capability length expected %d", capLen)) + } + + capBytes, err := core.ReadBytes(int(capLen)-4, r) + if err != nil { + return nil, err + } + capReader := bytes.NewReader(capBytes) + var c Capability + switch CapsType(capType) { + case CAPSTYPE_GENERAL: + c = &GeneralCapability{} + case CAPSTYPE_BITMAP: + c = &BitmapCapability{} + case CAPSTYPE_ORDER: + c = &OrderCapability{} + case CAPSTYPE_BITMAPCACHE: + c = &BitmapCacheCapability{} + case CAPSTYPE_POINTER: + c = &PointerCapability{} + case CAPSTYPE_INPUT: + c = &InputCapability{} + case CAPSTYPE_BRUSH: + c = &BrushCapability{} + case CAPSTYPE_GLYPHCACHE: + c = &GlyphCapability{} + case CAPSTYPE_OFFSCREENCACHE: + c = &OffscreenBitmapCacheCapability{} + case CAPSTYPE_VIRTUALCHANNEL: + c = &VirtualChannelCapability{} + // VCChunkSize is optional (MS-RDPBCGR 2.2.7.1.10); pad if absent + if len(capBytes) < 8 { + padded := make([]byte, 8) + copy(padded, capBytes) + capReader.Reset(padded) + } + case CAPSTYPE_SOUND: + c = &SoundCapability{} + case CAPSTYPE_CONTROL: + c = &ControlCapability{} + case CAPSTYPE_ACTIVATION: + c = &WindowActivationCapability{} + case CAPSTYPE_FONT: + c = &FontCapability{} + case CAPSTYPE_COLORCACHE: + c = &ColorCacheCapability{} + case CAPSTYPE_SHARE: + c = &ShareCapability{} + case CAPSETTYPE_MULTIFRAGMENTUPDATE: + c = &MultiFragmentUpdate{} + case CAPSTYPE_DRAWGDIPLUS: + c = &DrawGDIPlusCapability{} + case CAPSETTYPE_BITMAP_CODECS: + c = &BitmapCodecsCapability{} + case CAPSTYPE_BITMAPCACHE_HOSTSUPPORT: + c = &BitmapCacheHostSupportCapability{} + case CAPSETTYPE_LARGE_POINTER: + c = &LargePointerCapability{} + case CAPSTYPE_RAIL: + c = &RemoteProgramsCapability{} + case CAPSTYPE_WINDOW: + c = &WindowListCapability{} + case CAPSETTYPE_COMPDESK: + c = &DesktopCompositionCapability{} + case CAPSETTYPE_SURFACE_COMMANDS: + c = &SurfaceCommandsCapability{} + case CAPSSETTYPE_FRAME_ACKNOWLEDGE: + c = &FrameAcknowledgeCapability{} + default: + err := errors.New(fmt.Sprintf("unsupported Capability type 0x%04x", capType)) + slog.Error("readCapability", "err", err) + return nil, err + } + if err := struc.Unpack(capReader, c); err != nil { + slog.Error("readCapability", "err", err, "capType", capType, "capBytes", hex.EncodeToString(capBytes)) + return nil, err + } + slog.Debug("Capability", "type", c.Type(), "value", c) + return c, nil +} diff --git a/protocol/pdu/data.go b/protocol/pdu/data.go new file mode 100644 index 0000000..d796098 --- /dev/null +++ b/protocol/pdu/data.go @@ -0,0 +1,2337 @@ +package pdu + +import ( + "bytes" + "encoding/binary" + "errors" + "fmt" + "io" + "log/slog" + "sync" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/core" +) + +// capBuffPool pools bytes.Buffer instances used to serialize individual +// capability structures inside DemandActivePDU and ConfirmActivePDU. +// Each Serialize call borrows one buffer and returns it when done. +var capBuffPool = sync.Pool{ + New: func() any { return &bytes.Buffer{} }, +} + +// nscPlaneBufPool reuses byte slices for the intermediate YCoCg planes in +// decodeNSCodec. The planes are local to the decode call and released before +// the function returns, so pool re-use is safe. +var nscPlaneBufPool = sync.Pool{ + New: func() any { return []byte(nil) }, +} + +func acquireNSCPlaneBuf(size int) []byte { + b := nscPlaneBufPool.Get().([]byte) + if cap(b) >= size { + return b[:size] + } + return make([]byte, size) +} + +func releaseNSCPlaneBuf(b []byte) { + if b != nil { + nscPlaneBufPool.Put(b[:cap(b)]) + } +} + +// DecodeRemoteFX is a pluggable decoder for RemoteFX (MS-RDPRFX) surface codec +// data. It is set at init time by the main client package to avoid a circular +// import between protocol/pdu and plugin/rdpgfx. +// The function receives raw RFX data and returns top-down BGRA pixels. +var DecodeRemoteFX func(data []byte, width, height int) []byte + +const ( + PDUTYPE_DEMANDACTIVEPDU = 0x11 + PDUTYPE_CONFIRMACTIVEPDU = 0x13 + PDUTYPE_DEACTIVATEALLPDU = 0x16 + PDUTYPE_DATAPDU = 0x17 + PDUTYPE_SERVER_REDIR_PKT = 0x1A +) + +type PduType2 uint8 + +const ( + PDUTYPE2_UPDATE = 0x02 + PDUTYPE2_CONTROL = 0x14 + PDUTYPE2_POINTER = 0x1B + PDUTYPE2_INPUT = 0x1C + PDUTYPE2_SYNCHRONIZE = 0x1F + PDUTYPE2_REFRESH_RECT = 0x21 + PDUTYPE2_PLAY_SOUND = 0x22 + PDUTYPE2_SUPPRESS_OUTPUT = 0x23 + PDUTYPE2_SHUTDOWN_REQUEST = 0x24 + PDUTYPE2_SHUTDOWN_DENIED = 0x25 + PDUTYPE2_SAVE_SESSION_INFO = 0x26 + PDUTYPE2_FONTLIST = 0x27 + PDUTYPE2_FONTMAP = 0x28 + PDUTYPE2_SET_KEYBOARD_INDICATORS = 0x29 + PDUTYPE2_BITMAPCACHE_PERSISTENT_LIST = 0x2B + PDUTYPE2_BITMAPCACHE_ERROR_PDU = 0x2C + PDUTYPE2_SET_KEYBOARD_IME_STATUS = 0x2D + PDUTYPE2_OFFSCRCACHE_ERROR_PDU = 0x2E + PDUTYPE2_SET_ERROR_INFO_PDU = 0x2F + PDUTYPE2_DRAWNINEGRID_ERROR_PDU = 0x30 + PDUTYPE2_DRAWGDIPLUS_ERROR_PDU = 0x31 + PDUTYPE2_ARC_STATUS_PDU = 0x32 + PDUTYPE2_STATUS_INFO_PDU = 0x36 + PDUTYPE2_MONITOR_LAYOUT_PDU = 0x37 + PDUTYPE2_FRAME_ACKNOWLEDGE = 0x38 +) + +// Slow-Path Pointer Update types (MS-RDPBCGR 2.2.9.1.1.4) +const ( + TS_PTRUPDATE_TYPE_SYSTEM = 0x0001 + TS_PTRUPDATE_TYPE_POSITION = 0x0003 + TS_PTRUPDATE_TYPE_COLOR = 0x0006 + TS_PTRUPDATE_TYPE_CACHED = 0x0007 + TS_PTRUPDATE_TYPE_POINTER = 0x0008 +) + +func (p PduType2) String() string { + switch p { + case PDUTYPE2_UPDATE: + return "PDUTYPE2_UPDATE" + case PDUTYPE2_CONTROL: + return "PDUTYPE2_CONTROL" + case PDUTYPE2_POINTER: + return "PDUTYPE2_POINTER" + case PDUTYPE2_INPUT: + return "PDUTYPE2_INPUT" + case PDUTYPE2_SYNCHRONIZE: + return "PDUTYPE2_SYNCHRONIZE" + case PDUTYPE2_REFRESH_RECT: + return "PDUTYPE2_REFRESH_RECT" + case PDUTYPE2_PLAY_SOUND: + return "PDUTYPE2_PLAY_SOUND" + case PDUTYPE2_SUPPRESS_OUTPUT: + return "PDUTYPE2_SUPPRESS_OUTPUT" + case PDUTYPE2_SHUTDOWN_REQUEST: + return "PDUTYPE2_SHUTDOWN_REQUEST" + case PDUTYPE2_SHUTDOWN_DENIED: + return "PDUTYPE2_SHUTDOWN_DENIED" + case PDUTYPE2_SAVE_SESSION_INFO: + return "PDUTYPE2_SAVE_SESSION_INFO" + case PDUTYPE2_FONTLIST: + return "PDUTYPE2_FONTLIST" + case PDUTYPE2_FONTMAP: + return "PDUTYPE2_FONTMAP" + case PDUTYPE2_SET_KEYBOARD_INDICATORS: + return "PDUTYPE2_SET_KEYBOARD_INDICATORS" + case PDUTYPE2_BITMAPCACHE_PERSISTENT_LIST: + return "PDUTYPE2_BITMAPCACHE_PERSISTENT_LIST" + case PDUTYPE2_BITMAPCACHE_ERROR_PDU: + return "PDUTYPE2_BITMAPCACHE_ERROR_PDU" + case PDUTYPE2_SET_KEYBOARD_IME_STATUS: + return "PDUTYPE2_SET_KEYBOARD_IME_STATUS" + case PDUTYPE2_OFFSCRCACHE_ERROR_PDU: + return "PDUTYPE2_OFFSCRCACHE_ERROR_PDU" + case PDUTYPE2_SET_ERROR_INFO_PDU: + return "PDUTYPE2_SET_ERROR_INFO_PDU" + case PDUTYPE2_DRAWNINEGRID_ERROR_PDU: + return "PDUTYPE2_DRAWNINEGRID_ERROR_PDU" + case PDUTYPE2_DRAWGDIPLUS_ERROR_PDU: + return "PDUTYPE2_DRAWGDIPLUS_ERROR_PDU" + case PDUTYPE2_ARC_STATUS_PDU: + return "PDUTYPE2_ARC_STATUS_PDU" + case PDUTYPE2_STATUS_INFO_PDU: + return "PDUTYPE2_STATUS_INFO_PDU" + case PDUTYPE2_MONITOR_LAYOUT_PDU: + return "PDUTYPE2_MONITOR_LAYOUT_PDU" + } + + return "Unknown" +} + +const ( + CTRLACTION_REQUEST_CONTROL = 0x0001 + CTRLACTION_GRANTED_CONTROL = 0x0002 + CTRLACTION_DETACH = 0x0003 + CTRLACTION_COOPERATE = 0x0004 +) + +const ( + STREAM_UNDEFINED = 0x00 + STREAM_LOW = 0x01 + STREAM_MED = 0x02 + STREAM_HI = 0x04 +) + +type FastPathUpdateType uint8 + +const ( + FASTPATH_UPDATETYPE_ORDERS = 0x0 + FASTPATH_UPDATETYPE_BITMAP = 0x1 + FASTPATH_UPDATETYPE_PALETTE = 0x2 + FASTPATH_UPDATETYPE_SYNCHRONIZE = 0x3 + FASTPATH_UPDATETYPE_SURFCMDS = 0x4 + FASTPATH_UPDATETYPE_PTR_NULL = 0x5 + FASTPATH_UPDATETYPE_PTR_DEFAULT = 0x6 + FASTPATH_UPDATETYPE_PTR_POSITION = 0x8 + FASTPATH_UPDATETYPE_COLOR = 0x9 + FASTPATH_UPDATETYPE_CACHED = 0xA + FASTPATH_UPDATETYPE_POINTER = 0xB + FASTPATH_UPDATETYPE_LARGE_POINTER = 0xC +) + +func (t FastPathUpdateType) String() string { + switch t { + case FASTPATH_UPDATETYPE_ORDERS: + return "FASTPATH_UPDATETYPE_ORDERS" + case FASTPATH_UPDATETYPE_BITMAP: + return "FASTPATH_UPDATETYPE_BITMAP" + case FASTPATH_UPDATETYPE_PALETTE: + return "FASTPATH_UPDATETYPE_PALETTE" + case FASTPATH_UPDATETYPE_SYNCHRONIZE: + return "FASTPATH_UPDATETYPE_SYNCHRONIZE" + case FASTPATH_UPDATETYPE_SURFCMDS: + return "FASTPATH_UPDATETYPE_SURFCMDS" + case FASTPATH_UPDATETYPE_PTR_NULL: + return "FASTPATH_UPDATETYPE_PTR_NULL" + case FASTPATH_UPDATETYPE_PTR_DEFAULT: + return "FASTPATH_UPDATETYPE_PTR_DEFAULT" + case FASTPATH_UPDATETYPE_PTR_POSITION: + return "FASTPATH_UPDATETYPE_PTR_POSITION" + case FASTPATH_UPDATETYPE_COLOR: + return "FASTPATH_UPDATETYPE_COLOR" + case FASTPATH_UPDATETYPE_CACHED: + return "FASTPATH_UPDATETYPE_CACHED" + case FASTPATH_UPDATETYPE_POINTER: + return "FASTPATH_UPDATETYPE_POINTER" + case FASTPATH_UPDATETYPE_LARGE_POINTER: + return "FASTPATH_UPDATETYPE_LARGE_POINTER" + } + + return "Unknown" +} + +const ( + BITMAP_COMPRESSION = 0x0001 + //NO_BITMAP_COMPRESSION_HDR = 0x0400 + BITMAP_NO_PROCESSING = 0x8000 // Surface command: data is already decoded top-down BGRA +) + +// Surface Command types (MS-RDPBCGR 2.2.9.1.2.1) +const ( + CMDTYPE_SET_SURFACE_BITS = 0x0001 + CMDTYPE_FRAME_MARKER = 0x0004 + CMDTYPE_STREAM_SURFACE_BITS = 0x0006 +) + +const ( + SURFCMD_FRAMEACTION_BEGIN = 0x0000 + SURFCMD_FRAMEACTION_END = 0x0001 +) + +/* compression types */ +const ( + RDP_MPPC_BIG = 0x01 + RDP_MPPC_COMPRESSED = 0x20 + RDP_MPPC_RESET = 0x40 + RDP_MPPC_FLUSH = 0x80 + RDP_MPPC_DICT_SIZE = 65536 +) + +type ShareDataHeader struct { + SharedId uint32 `struc:"little"` + Padding1 uint8 `struc:"little"` + StreamId uint8 `struc:"little"` + UncompressedLength uint16 `struc:"little"` + PDUType2 uint8 `struc:"little"` + CompressedType uint8 `struc:"little"` + CompressedLength uint16 `struc:"little"` +} + +func NewShareDataHeader(size int, type2 uint8, shareId uint32) *ShareDataHeader { + return &ShareDataHeader{ + SharedId: shareId, + PDUType2: type2, + StreamId: STREAM_LOW, + UncompressedLength: uint16(size + 4), + } +} + +type PDUMessage interface { + Type() uint16 + Serialize() []byte +} + +type DemandActivePDU struct { + SharedId uint32 `struc:"little"` + LengthSourceDescriptor uint16 `struc:"little,sizeof=SourceDescriptor"` + LengthCombinedCapabilities uint16 `struc:"little"` + SourceDescriptor []byte `struc:"sizefrom=LengthSourceDescriptor"` + NumberCapabilities uint16 `struc:"little,sizeof=CapabilitySets"` + Pad2Octets uint16 `struc:"little"` + CapabilitySets []Capability `struc:"sizefrom=NumberCapabilities"` + SessionId uint32 `struc:"little"` +} + +func (d *DemandActivePDU) Type() uint16 { + return PDUTYPE_DEMANDACTIVEPDU +} + +func (d *DemandActivePDU) Serialize() []byte { + buff := &bytes.Buffer{} + core.WriteUInt32LE(d.SharedId, buff) + core.WriteUInt16LE(d.LengthSourceDescriptor, buff) + core.WriteUInt16LE(d.LengthCombinedCapabilities, buff) + core.WriteBytes([]byte(d.SourceDescriptor), buff) + core.WriteUInt16LE(uint16(len(d.CapabilitySets)), buff) + core.WriteUInt16LE(d.Pad2Octets, buff) + capBuff := capBuffPool.Get().(*bytes.Buffer) + for _, cap := range d.CapabilitySets { + core.WriteUInt16LE(uint16(cap.Type()), buff) + capBuff.Reset() + struc.Pack(capBuff, cap) + capBytes := capBuff.Bytes() + core.WriteUInt16LE(uint16(len(capBytes)+4), buff) + core.WriteBytes(capBytes, buff) + } + capBuffPool.Put(capBuff) + core.WriteUInt32LE(d.SessionId, buff) + return buff.Bytes() +} + +func readDemandActivePDU(r io.Reader) (*DemandActivePDU, error) { + d := &DemandActivePDU{} + var err error + d.SharedId, err = core.ReadUInt32LE(r) + if err != nil { + return nil, err + } + d.LengthSourceDescriptor, err = core.ReadUint16LE(r) + d.LengthCombinedCapabilities, err = core.ReadUint16LE(r) + sourceDescriptorBytes, err := core.ReadBytes(int(d.LengthSourceDescriptor), r) + if err != nil { + return nil, err + } + d.SourceDescriptor = sourceDescriptorBytes + d.NumberCapabilities, err = core.ReadUint16LE(r) + d.Pad2Octets, err = core.ReadUint16LE(r) + d.CapabilitySets = make([]Capability, 0, d.NumberCapabilities) + for i := 0; i < int(d.NumberCapabilities); i++ { + c, err := readCapability(r) + if err != nil { + //return nil, err + continue + } + d.CapabilitySets = append(d.CapabilitySets, c) + } + d.NumberCapabilities = uint16(len(d.CapabilitySets)) + d.SessionId, err = core.ReadUInt32LE(r) + if err != nil { + return nil, err + } + return d, nil +} + +type ConfirmActivePDU struct { + SharedId uint32 `struc:"little"` + OriginatorId uint16 `struc:"little"` + LengthSourceDescriptor uint16 `struc:"little,sizeof=SourceDescriptor"` + LengthCombinedCapabilities uint16 `struc:"little"` + SourceDescriptor []byte `struc:"sizefrom=LengthSourceDescriptor"` + NumberCapabilities uint16 `struc:"little,sizeof=CapabilitySets"` + Pad2Octets uint16 `struc:"little"` + CapabilitySets []Capability `struc:"sizefrom=NumberCapabilities"` +} + +func (*ConfirmActivePDU) Type() uint16 { + return PDUTYPE_CONFIRMACTIVEPDU +} + +func (c *ConfirmActivePDU) Serialize() []byte { + buff := &bytes.Buffer{} + core.WriteUInt32LE(c.SharedId, buff) + core.WriteUInt16LE(c.OriginatorId, buff) + core.WriteUInt16LE(uint16(len(c.SourceDescriptor)), buff) + + capsBuff := &bytes.Buffer{} + capBuff := capBuffPool.Get().(*bytes.Buffer) + for _, capa := range c.CapabilitySets { + core.WriteUInt16LE(uint16(capa.Type()), capsBuff) + capBuff.Reset() + struc.Pack(capBuff, capa) + capBytes := capBuff.Bytes() + core.WriteUInt16LE(uint16(len(capBytes)+4), capsBuff) + core.WriteBytes(capBytes, capsBuff) + } + capBuffPool.Put(capBuff) + capsBytes := capsBuff.Bytes() + + core.WriteUInt16LE(uint16(2+2+len(capsBytes)), buff) + core.WriteBytes(c.SourceDescriptor, buff) + core.WriteUInt16LE(c.NumberCapabilities, buff) + core.WriteUInt16LE(c.Pad2Octets, buff) + core.WriteBytes(capsBytes, buff) + return buff.Bytes() +} + +// 9401 => share control header +// 1300 => share control header +// ec03 => share control header +// ea030100 => shareId 66538 +// ea03 => OriginatorId +// 0400 +// 8001 => LengthCombinedCapabilities +// 72647079 +// 0c00 => NumberCapabilities 12 +// 0000 +// caps below +// 010018000100030000020000000015040000000000000000 +// 02001c00180001000100010000052003000000000100000001000000 +// 030058000000000000000000000000000000000000000000010014000000010000000a0000000000000000000000000000000000000000000000000000000000000000000000000000000000008403000000000000000000 +// 04002800000000000000000000000000000000000000000000000000000000000000000000000000 +// 0800080000001400 +// 0c00080000000000 +// 0d005c001500000009040000040000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000c000000 +// 0f00080000000000 +// 10003400000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000 +// 11000c000000000000000000 +// 14000c000000000000000000 +// 1a00080000000000 + +func NewConfirmActivePDU() *ConfirmActivePDU { + return &ConfirmActivePDU{ + OriginatorId: 0x03EA, + CapabilitySets: make([]Capability, 0), + SourceDescriptor: []byte("rdpy"), + } +} + +func readConfirmActivePDU(r io.Reader) (*ConfirmActivePDU, error) { + p := &ConfirmActivePDU{} + var err error + p.SharedId, err = core.ReadUInt32LE(r) + if err != nil { + return nil, err + } + p.OriginatorId, err = core.ReadUint16LE(r) + p.LengthSourceDescriptor, err = core.ReadUint16LE(r) + p.LengthCombinedCapabilities, err = core.ReadUint16LE(r) + + sourceDescriptorBytes, err := core.ReadBytes(int(p.LengthSourceDescriptor), r) + if err != nil { + return nil, err + } + p.SourceDescriptor = sourceDescriptorBytes + p.NumberCapabilities, err = core.ReadUint16LE(r) + p.Pad2Octets, err = core.ReadUint16LE(r) + + p.CapabilitySets = make([]Capability, 0, p.NumberCapabilities) + for i := 0; i < int(p.NumberCapabilities); i++ { + c, err := readCapability(r) + if err != nil { + return nil, err + } + p.CapabilitySets = append(p.CapabilitySets, c) + } + s, _ := core.ReadUInt32LE(r) + slog.Debug("readConfirmActivePDU", "sessionid", s) + return p, nil +} + +type DeactiveAllPDU struct { + ShareId uint32 `struc:"little"` + LengthSourceDescriptor uint16 `struc:"little,sizeof=SourceDescriptor"` + SourceDescriptor []byte +} + +func (*DeactiveAllPDU) Type() uint16 { + return PDUTYPE_DEACTIVATEALLPDU +} + +func (d *DeactiveAllPDU) Serialize() []byte { + buff := &bytes.Buffer{} + struc.Pack(buff, d) + return buff.Bytes() +} + +func readDeactiveAllPDU(r io.Reader) (*DeactiveAllPDU, error) { + p := &DeactiveAllPDU{} + err := struc.Unpack(r, p) + return p, err +} + +// ServerRedirectionPDU represents the RDP Server Redirection PDU +// (MS-RDPBCGR 2.2.13.2.1). Only the LoadBalanceInfo field (routing +// token) is extracted; other optional fields are skipped. +type ServerRedirectionPDU struct { + Flags uint16 + Length uint16 + SessionID uint32 + RedirFlags uint32 + LoadBalanceInfo []byte +} + +const ( + LB_TARGET_NET_ADDRESS = 0x00000001 + LB_LOAD_BALANCE_INFO = 0x00000002 + LB_USERNAME = 0x00000004 +) + +func (*ServerRedirectionPDU) Type() uint16 { + return PDUTYPE_SERVER_REDIR_PKT +} + +func (d *ServerRedirectionPDU) Serialize() []byte { + return nil +} + +func readServerRedirectionPDU(r io.Reader) (*ServerRedirectionPDU, error) { + // Enhanced Security variant has a 2-byte pad before the PDU body + if _, err := core.ReadUint16LE(r); err != nil { + return nil, fmt.Errorf("redir: read pad: %w", err) + } + + redir := &ServerRedirectionPDU{} + var err error + if redir.Flags, err = core.ReadUint16LE(r); err != nil { + return nil, fmt.Errorf("redir: read flags: %w", err) + } + if redir.Length, err = core.ReadUint16LE(r); err != nil { + return nil, fmt.Errorf("redir: read length: %w", err) + } + if redir.SessionID, err = core.ReadUInt32LE(r); err != nil { + return nil, fmt.Errorf("redir: read sessionID: %w", err) + } + if redir.RedirFlags, err = core.ReadUInt32LE(r); err != nil { + return nil, fmt.Errorf("redir: read redirFlags: %w", err) + } + + // Parse variable-length fields in flag order. + // We only need LoadBalanceInfo (routing token) for reconnection. + if redir.RedirFlags&LB_TARGET_NET_ADDRESS != 0 { + cbLen, err := core.ReadUInt32LE(r) + if err != nil { + return nil, fmt.Errorf("redir: read targetNetAddr len: %w", err) + } + if _, err := core.ReadBytes(int(cbLen), r); err != nil { + return nil, fmt.Errorf("redir: read targetNetAddr: %w", err) + } + } + + if redir.RedirFlags&LB_LOAD_BALANCE_INFO != 0 { + cbLen, err := core.ReadUInt32LE(r) + if err != nil { + return nil, fmt.Errorf("redir: read loadBalanceInfo len: %w", err) + } + redir.LoadBalanceInfo, err = core.ReadBytes(int(cbLen), r) + if err != nil { + return nil, fmt.Errorf("redir: read loadBalanceInfo: %w", err) + } + } + + slog.Debug("Server Redirection PDU", + "flags", redir.Flags, + "sessionID", redir.SessionID, + "redirFlags", redir.RedirFlags, + "loadBalanceInfo", string(redir.LoadBalanceInfo)) + return redir, nil +} + +type DataPDU struct { + Header *ShareDataHeader + Data DataPDUData +} + +func (*DataPDU) Type() uint16 { + return PDUTYPE_DATAPDU +} + +func (d *DataPDU) Serialize() []byte { + buff := &bytes.Buffer{} + struc.Pack(buff, d.Header) + struc.Pack(buff, d.Data) + return buff.Bytes() +} + +func NewDataPDU(data DataPDUData, shareId uint32) *DataPDU { + dataLen, err := struc.Sizeof(data) + if err != nil { + // Fallback: pack to measure length + dataBuff := &bytes.Buffer{} + struc.Pack(dataBuff, data) + dataLen = dataBuff.Len() + } + return &DataPDU{ + Header: NewShareDataHeader(dataLen, data.Type2(), shareId), + Data: data, + } +} + +func readDataPDU(r io.Reader, mppc *core.MppcDecompressor) (*DataPDU, error) { + header := &ShareDataHeader{} + err := struc.Unpack(r, header) + if err != nil { + slog.Error("readDataPDU", "err", err) + return nil, err + } + + // Decompress the payload when the server has compressed it. + if header.CompressedType != 0 && mppc != nil { + if header.CompressedType&RDP_MPPC_COMPRESSED != 0 { + compressed, err := core.ReadBytes(int(header.CompressedLength), r) + if err != nil { + slog.Error("readDataPDU: reading compressed payload", "err", err) + return nil, err + } + decompressed, err := mppc.Decompress(header.CompressedType, compressed) + if err != nil { + slog.Error("readDataPDU: MPPC decompression failed", "err", err) + return nil, err + } + r = bytes.NewReader(decompressed) + } else { + // CompressedType set but not COMPRESSED: update history only. + plain, _ := core.ReadBytes(int(header.CompressedLength), r) + _, _ = mppc.Decompress(header.CompressedType, plain) + r = bytes.NewReader(plain) + } + } + + var d DataPDUData + slog.Debug("readDataPDU", "PDUTYPE2", header.PDUType2) + switch header.PDUType2 { + case PDUTYPE2_UPDATE: + d = &UpdateDataPDU{} + + case PDUTYPE2_SYNCHRONIZE: + d = &SynchronizeDataPDU{} + + case PDUTYPE2_CONTROL: + d = &ControlDataPDU{} + + case PDUTYPE2_FONTLIST: + d = &FontListDataPDU{} + + case PDUTYPE2_SET_ERROR_INFO_PDU: + d = &ErrorInfoDataPDU{} + + case PDUTYPE2_FONTMAP: + d = &FontMapDataPDU{} + + case PDUTYPE2_SAVE_SESSION_INFO: + d = &SaveSessionInfo{} + + case PDUTYPE2_POINTER: + d = &PointerDataPDU{} + + case PDUTYPE2_SET_KEYBOARD_INDICATORS: + d = &SetKeyboardIndicatorsDataPDU{} + + default: + err = fmt.Errorf("Unknown data pdu type2 0x%02x", header.PDUType2) + slog.Error("readDataPDU", "err", err) + return nil, err + } + + err = d.Unpack(r) + if err != nil { + slog.Error("readDataPDU", "err", err) + return nil, err + } + + p := &DataPDU{ + Header: header, + Data: d, + } + return p, nil +} + +type DataPDUData interface { + Type2() uint8 + Unpack(io.Reader) error +} + +type UpdateDataPDU struct { + UpdateType uint16 + Udata UpdateData +} + +func (*UpdateDataPDU) Type2() uint8 { + return PDUTYPE2_UPDATE +} +func (d *UpdateDataPDU) Unpack(r io.Reader) (err error) { + //slow path update + d.UpdateType, err = core.ReadUint16LE(r) + slog.Debug("FastPathUpdate", "type", d.UpdateType) + var p UpdateData + switch d.UpdateType { + case FASTPATH_UPDATETYPE_ORDERS: + case FASTPATH_UPDATETYPE_BITMAP: + p = &BitmapUpdateDataPDU{} + case FASTPATH_UPDATETYPE_PALETTE: + case FASTPATH_UPDATETYPE_SYNCHRONIZE: + } + if p != nil { + err = p.Unpack(r) + if err != nil { + //slog.Error("Unpack:", err) + return err + } + } else { + return fmt.Errorf("Unsupport slow update type 0x%x", d.UpdateType) + } + + d.Udata = p + + return nil +} + +// PointerDataPDU handles slow-path pointer updates (MS-RDPBCGR 2.2.9.1.1.4) +type PointerDataPDU struct { + MessageType uint16 + Pad2Octets uint16 + Pdata UpdateData +} + +func (*PointerDataPDU) Type2() uint8 { + return PDUTYPE2_POINTER +} + +func (d *PointerDataPDU) Unpack(r io.Reader) error { + var err error + d.MessageType, err = core.ReadUint16LE(r) + if err != nil { + return err + } + d.Pad2Octets, err = core.ReadUint16LE(r) + if err != nil { + return err + } + slog.Debug("PointerDataPDU", "messageType", d.MessageType) + var p UpdateData + switch d.MessageType { + case TS_PTRUPDATE_TYPE_CACHED: + p = &FastPathUpdateCachedPDU{} + case TS_PTRUPDATE_TYPE_POINTER: + p = &FastPathUpdatePointerPDU{} + case TS_PTRUPDATE_TYPE_SYSTEM, TS_PTRUPDATE_TYPE_POSITION, TS_PTRUPDATE_TYPE_COLOR: + // not yet parsed; remaining data is discarded by the caller + default: + slog.Debug("PointerDataPDU: unhandled", "messageType", d.MessageType) + } + if p != nil { + if err = p.Unpack(r); err != nil { + return err + } + } + d.Pdata = p + return nil +} + +func (d *PointerDataPDU) Serialize() []byte { + return nil +} + +type BitmapUpdateDataPDU struct { + NumberRectangles uint16 `struc:"little,sizeof=Rectangles"` + Rectangles []BitmapData +} + +func (*BitmapUpdateDataPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_BITMAP +} +func (f *BitmapUpdateDataPDU) Unpack(r io.Reader) error { + var err error + f.NumberRectangles, err = core.ReadUint16LE(r) + if err != nil { + return err + } + if f.NumberRectangles > 4096 { + return fmt.Errorf("implausible rectangle count %d", f.NumberRectangles) + } + f.Rectangles = make([]BitmapData, 0, f.NumberRectangles) + for i := 0; i < int(f.NumberRectangles); i++ { + rect := BitmapData{} + rect.DestLeft, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestTop, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestRight, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestBottom, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Width, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Height, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitsPerPixel, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Flags, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapLength, err = core.ReadUint16LE(r) + if err != nil { + return err + } + ln := rect.BitmapLength + if rect.Flags&BITMAP_COMPRESSION != 0 && (rect.Flags&NO_BITMAP_COMPRESSION_HDR == 0) { + rect.BitmapComprHdr = new(BitmapCompressedDataHeader) + rect.BitmapComprHdr.CbCompFirstRowSize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbCompMainBodySize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbScanWidth, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbUncompressedSize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + ln = rect.BitmapComprHdr.CbCompMainBodySize + } + + rect.BitmapDataStream, err = core.ReadBytes(int(ln), r) + if err != nil { + return err + } + f.Rectangles = append(f.Rectangles, rect) + } + return nil +} + +type SynchronizeDataPDU struct { + MessageType uint16 `struc:"little"` + TargetUser uint16 `struc:"little"` +} + +func (*SynchronizeDataPDU) Type2() uint8 { + return PDUTYPE2_SYNCHRONIZE +} + +func NewSynchronizeDataPDU(targetUser uint16) *SynchronizeDataPDU { + return &SynchronizeDataPDU{ + MessageType: 1, + TargetUser: targetUser, + } +} +func (d *SynchronizeDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type ControlDataPDU struct { + Action uint16 `struc:"little"` + GrantId uint16 `struc:"little"` + ControlId uint32 `struc:"little"` +} + +func (*ControlDataPDU) Type2() uint8 { + return PDUTYPE2_CONTROL +} +func (d *ControlDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type FontListDataPDU struct { + NumberFonts uint16 `struc:"little"` + TotalNumFonts uint16 `struc:"little"` + ListFlags uint16 `struc:"little"` + EntrySize uint16 `struc:"little"` +} + +func (*FontListDataPDU) Type2() uint8 { + return PDUTYPE2_FONTLIST +} +func (d *FontListDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type ErrorInfoDataPDU struct { + ErrorInfo uint32 `struc:"little"` +} + +func (*ErrorInfoDataPDU) Type2() uint8 { + return PDUTYPE2_SET_ERROR_INFO_PDU +} +func (d *ErrorInfoDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type FontMapDataPDU struct { + NumberEntries uint16 `struc:"little"` + TotalNumEntries uint16 `struc:"little"` + MapFlags uint16 `struc:"little"` + EntrySize uint16 `struc:"little"` +} + +func (*FontMapDataPDU) Type2() uint8 { + return PDUTYPE2_FONTMAP +} +func (d *FontMapDataPDU) Unpack(r io.Reader) error { + err := struc.Unpack(r, d) + // MS-RDPBCGR 2.2.1.22.1: Font Map payload fields are optional. + // VirtualBox sends a short FontMap PDU with no payload data. + if err == io.EOF || err == io.ErrUnexpectedEOF { + return nil + } + return err +} + +// SetKeyboardIndicatorsDataPDU sets the state of keyboard indicator LEDs. +// MS-RDPBCGR 2.2.8.2.1.3.3.1 +type SetKeyboardIndicatorsDataPDU struct { + UnitId uint16 `struc:"little"` + LedFlags uint16 `struc:"little"` +} + +func (*SetKeyboardIndicatorsDataPDU) Type2() uint8 { + return PDUTYPE2_SET_KEYBOARD_INDICATORS +} +func (d *SetKeyboardIndicatorsDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +// SuppressOutputPDU tells the server to start/stop sending display updates. +// MS-RDPBCGR 2.2.11.3.1 +type SuppressOutputPDU struct { + AllowDisplayUpdates uint8 `struc:"little"` + Pad3Octets [3]byte `struc:"little"` + Left uint16 `struc:"little"` + Top uint16 `struc:"little"` + Right uint16 `struc:"little"` + Bottom uint16 `struc:"little"` +} + +func (*SuppressOutputPDU) Type2() uint8 { + return PDUTYPE2_SUPPRESS_OUTPUT +} +func (d *SuppressOutputPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +// FrameAcknowledgeDataPDU acknowledges receipt of a frame (MS-RDPBCGR 2.2.11.3.2). +type FrameAcknowledgeDataPDU struct { + FrameID uint32 `struc:"little"` +} + +func (*FrameAcknowledgeDataPDU) Type2() uint8 { + return PDUTYPE2_FRAME_ACKNOWLEDGE +} +func (d *FrameAcknowledgeDataPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +// RefreshRectPDU requests the server to redraw one or more screen regions. +// MS-RDPBCGR 2.2.11.2 +type RefreshRectPDU struct { + NumberOfAreas uint8 `struc:"little"` + Pad3Octets [3]byte `struc:"little"` + Left uint16 `struc:"little"` + Top uint16 `struc:"little"` + Right uint16 `struc:"little"` + Bottom uint16 `struc:"little"` +} + +func (*RefreshRectPDU) Type2() uint8 { + return PDUTYPE2_REFRESH_RECT +} +func (d *RefreshRectPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type InfoType uint32 + +const ( + INFOTYPE_LOGON = 0x00000000 + INFOTYPE_LOGON_LONG = 0x00000001 + INFOTYPE_LOGON_PLAINNOTIFY = 0x00000002 + INFOTYPE_LOGON_EXTENDED_INFO = 0x00000003 +) +const ( + LOGON_EX_AUTORECONNECTCOOKIE = 0x00000001 + LOGON_EX_LOGONERRORS = 0x00000002 +) + +type LogonFields struct { + CbFileData uint32 `struc:"little"` + Len uint32 //28 `struc:"little"` + Version uint32 // 1 `struc:"little"` + LogonId uint32 `struc:"little"` + random [16]byte //16 `struc:"little"` +} +type SaveSessionInfo struct { + InfoType uint32 + Length uint16 + FieldsPresent uint32 + LogonId uint32 + Random []byte +} + +func (s *SaveSessionInfo) logonInfoV1(r io.Reader) (err error) { + core.ReadUInt32LE(r) // cbDomain + b, _ := core.ReadBytes(52, r) + domain := core.UnicodeDecode(b) + + core.ReadUInt32LE(r) // cbUserName + b, _ = core.ReadBytes(512, r) + userName := core.UnicodeDecode(b) + + sessionId, _ := core.ReadUInt32LE(r) + s.LogonId = sessionId + slog.Debug("logonInfo", "sessionId", s.LogonId, "userName", userName, "domain", domain) + return err +} +func (s *SaveSessionInfo) logonInfoV2(r io.Reader) (err error) { + core.ReadUint16LE(r) + core.ReadUInt32LE(r) + sessionId, _ := core.ReadUInt32LE(r) + s.LogonId = sessionId + cbDomain, _ := core.ReadUInt32LE(r) + cbUserName, _ := core.ReadUInt32LE(r) + core.ReadBytes(558, r) + + b, _ := core.ReadBytes(int(cbDomain), r) + domain := core.UnicodeDecode(b) + b, _ = core.ReadBytes(int(cbUserName), r) + userName := core.UnicodeDecode(b) + slog.Debug("logonInfoV2", "sessionId", s.LogonId, "userName", userName, "domain", domain) + + return err +} +func (s *SaveSessionInfo) logonPlainNotify(r io.Reader) (err error) { + core.ReadBytes(576, r) /* pad (576 bytes) */ + return err +} +func (s *SaveSessionInfo) logonInfoExtended(r io.Reader) (err error) { + s.Length, err = core.ReadUint16LE(r) + s.FieldsPresent, err = core.ReadUInt32LE(r) + //slog.Debug("FieldsPresent:", s.FieldsPresent) + // auto reconnect cookie + if s.FieldsPresent&LOGON_EX_AUTORECONNECTCOOKIE != 0 { + core.ReadUInt32LE(r) + b, _ := core.ReadUInt32LE(r) + if b != 28 { + return errors.New("invalid length in Auto-Reconnect packet") + } + b, _ = core.ReadUInt32LE(r) + if b != 1 { + return errors.New("unsupported version of Auto-Reconnect packet") + } + b, _ = core.ReadUInt32LE(r) + s.LogonId = b + s.Random, _ = core.ReadBytes(16, r) + } else { // logon error info + core.ReadUInt32LE(r) + b, _ := core.ReadUInt32LE(r) + b, _ = core.ReadUInt32LE(r) + s.LogonId = b + } + core.ReadBytes(570, r) + return err +} +func (s *SaveSessionInfo) Unpack(r io.Reader) (err error) { + s.InfoType, err = core.ReadUInt32LE(r) + switch s.InfoType { + case INFOTYPE_LOGON: + err = s.logonInfoV1(r) + case INFOTYPE_LOGON_LONG: + err = s.logonInfoV2(r) + case INFOTYPE_LOGON_PLAINNOTIFY: + err = s.logonPlainNotify(r) + case INFOTYPE_LOGON_EXTENDED_INFO: + err = s.logonInfoExtended(r) + default: + return fmt.Errorf("Unhandled saveSessionInfo type 0x%x", s.InfoType) + } + + return err +} + +func (*SaveSessionInfo) Type2() uint8 { + return PDUTYPE2_SAVE_SESSION_INFO +} + +type PersistKeyPDU struct { + NumEntriesCache0 uint16 `struc:"little"` + NumEntriesCache1 uint16 `struc:"little"` + NumEntriesCache2 uint16 `struc:"little"` + NumEntriesCache3 uint16 `struc:"little"` + NumEntriesCache4 uint16 `struc:"little"` + TotalEntriesCache0 uint16 `struc:"little"` + TotalEntriesCache1 uint16 `struc:"little"` + TotalEntriesCache2 uint16 `struc:"little"` + TotalEntriesCache3 uint16 `struc:"little"` + TotalEntriesCache4 uint16 `struc:"little"` + BBitMask uint8 `struc:"little"` + Pad1 uint8 `struc:"little"` + Ppad3 uint16 `struc:"little"` +} + +func (*PersistKeyPDU) Type2() uint8 { + return PDUTYPE2_BITMAPCACHE_PERSISTENT_LIST +} + +type UpdateData interface { + FastPathUpdateType() uint8 + Unpack(io.Reader) error +} + +type BitmapCompressedDataHeader struct { + CbCompFirstRowSize uint16 `struc:"little"` + CbCompMainBodySize uint16 `struc:"little"` + CbScanWidth uint16 `struc:"little"` + CbUncompressedSize uint16 `struc:"little"` +} + +type BitmapData struct { + DestLeft uint16 `struc:"little"` + DestTop uint16 `struc:"little"` + DestRight uint16 `struc:"little"` + DestBottom uint16 `struc:"little"` + Width uint16 `struc:"little"` + Height uint16 `struc:"little"` + BitsPerPixel uint16 `struc:"little"` + Flags uint16 `struc:"little"` + BitmapLength uint16 `struc:"little,sizeof=BitmapDataStream"` + BitmapComprHdr *BitmapCompressedDataHeader + BitmapDataStream []byte +} + +func (b *BitmapData) IsCompress() bool { + return b.Flags&BITMAP_COMPRESSION != 0 +} + +type FastPathBitmapUpdateDataPDU struct { + Header uint16 `struc:"little"` + NumberRectangles uint16 `struc:"little,sizeof=Rectangles"` + Rectangles []BitmapData +} + +func (f *FastPathBitmapUpdateDataPDU) Unpack(r io.Reader) error { + var err error + f.Header, err = core.ReadUint16LE(r) + if err != nil { + return err + } + f.NumberRectangles, err = core.ReadUint16LE(r) + if err != nil { + return err + } + // 矩形数受载荷物理限制(每矩形至少 18 字节);超限即流已错位, + // 直接拒绝而不是按垃圾矩形继续解析。 + if f.NumberRectangles > 4096 { + return fmt.Errorf("implausible rectangle count %d", f.NumberRectangles) + } + f.Rectangles = make([]BitmapData, 0, f.NumberRectangles) + for i := 0; i < int(f.NumberRectangles); i++ { + rect := BitmapData{} + rect.DestLeft, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestTop, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestRight, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.DestBottom, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Width, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Height, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitsPerPixel, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.Flags, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapLength, err = core.ReadUint16LE(r) + if err != nil { + return err + } + ln := rect.BitmapLength + if rect.Flags&BITMAP_COMPRESSION != 0 && (rect.Flags&NO_BITMAP_COMPRESSION_HDR == 0) { + rect.BitmapComprHdr = new(BitmapCompressedDataHeader) + rect.BitmapComprHdr.CbCompFirstRowSize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbCompMainBodySize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbScanWidth, err = core.ReadUint16LE(r) + if err != nil { + return err + } + rect.BitmapComprHdr.CbUncompressedSize, err = core.ReadUint16LE(r) + if err != nil { + return err + } + ln = rect.BitmapComprHdr.CbCompMainBodySize + } + + rect.BitmapDataStream, err = core.ReadBytes(int(ln), r) + if err != nil { + return err + } + f.Rectangles = append(f.Rectangles, rect) + } + return nil +} + +func (*FastPathBitmapUpdateDataPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_BITMAP +} + +type FastPathColorPdu struct { + CacheIdx uint16 + X uint16 + Y uint16 + Width uint16 + Height uint16 + MaskLen uint16 `struc:"little,sizeof=Mask"` + DataLen uint16 `struc:"little,sizeof=Data"` + Mask []byte + Data []byte +} + +func (*FastPathColorPdu) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_COLOR +} +func (f *FastPathColorPdu) Unpack(r io.Reader) error { + return struc.Unpack(r, f) +} + +type FastPathSurfaceCmds struct { + Rects []BitmapData +} + +func (*FastPathSurfaceCmds) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_SURFCMDS +} +func (f *FastPathSurfaceCmds) Unpack(r io.Reader) error { + // This won't be called; Surface Commands are handled directly in RecvFastPath. + return nil +} + +// SurfaceCommandsResult holds parsed bitmap data and frame IDs to acknowledge. +type SurfaceCommandsResult struct { + Rects []BitmapData + FrameIDs []uint32 +} + +// ParseSurfaceCommands parses one or more surface commands from raw data +// and returns decoded BitmapData rectangles and frame IDs that need acknowledgment. +func ParseSurfaceCommands(data []byte) SurfaceCommandsResult { + r := bytes.NewReader(data) + var result SurfaceCommandsResult + for r.Len() > 0 { + cmdType, err := core.ReadUint16LE(r) + if err != nil { + break + } + switch cmdType { + case CMDTYPE_SET_SURFACE_BITS, CMDTYPE_STREAM_SURFACE_BITS: + rect, err := decodeSurfaceBitsCmd(r, true) + if err != nil { + slog.Warn("decodeSurfaceBitsCmd", "err", err) + return result + } + if rect != nil { + result.Rects = append(result.Rects, *rect) + } + case CMDTYPE_FRAME_MARKER: + frameAction, _ := core.ReadUint16LE(r) + frameId, _ := core.ReadUInt32LE(r) + if frameAction == SURFCMD_FRAMEACTION_END { + result.FrameIDs = append(result.FrameIDs, frameId) + } + default: + slog.Warn("Unknown surface command type", "cmdType", cmdType) + return result + } + } + return result +} + +// decodeSurfaceBitsCmd parses a SET_SURFACE_BITS or STREAM_SURFACE_BITS command. +// win10Layout=true 时按 Win10 实测 22 字节头解析(codecID 前多一个保留字节), +// 否则按 FreeRDP 经典 21 字节头解析。命令不含开头的 cmdType(2)。 +func decodeSurfaceBitsCmd(r io.Reader, win10Layout bool) (*BitmapData, error) { + destLeft, err := core.ReadUint16LE(r) + if err != nil { + return nil, err + } + destTop, _ := core.ReadUint16LE(r) + destRight, _ := core.ReadUint16LE(r) + destBottom, _ := core.ReadUint16LE(r) + + // 位图头尾部有两种实测布局(均从命令起始计偏移): + // Win10 19041 实测(头长 22):bpp@10 ?@11 ?@12 codecID@13 width@14 + // height@16 bitmapDataLength u32@18,数据@22; + // FreeRDP update_recv_surface_bits(头长 21):bpp@10 reserved@11 + // codecID@12 width@13 height@15 bitmapDataLength u32@17,数据@21。 + // win10Layout 由调用方依据 width/height 与 dest 矩形的一致性判别, + // 相对经典布局在 codecID 前多 1 个保留字节。 + bpp, _ := core.ReadUInt8(r) + _, _ = core.ReadUInt8(r) // reserved + if win10Layout { + _, _ = core.ReadUInt8(r) + } + codecID, _ := core.ReadUInt8(r) + width, _ := core.ReadUint16LE(r) + height, _ := core.ReadUint16LE(r) + bitmapDataLength, _ := core.ReadUInt32LE(r) + + bitmapData, err := core.ReadBytes(int(bitmapDataLength), r) + if err != nil { + return nil, fmt.Errorf("failed to read bitmap data: %v", err) + } + + slog.Debug("decodeSurfaceBitsCmd", + "destLeft", destLeft, "destTop", destTop, "destRight", destRight, "destBottom", destBottom, + "width", width, "height", height, + "bpp", bpp, "codecID", codecID, "dataLen", bitmapDataLength) + + var pixels []byte + outBpp := uint16(bpp) + switch codecID { + case 0: // Uncompressed + pixels = bitmapData + case 1: // NSCodec + pixels = decodeNSCodec(bitmapData, int(width), int(height)) + outBpp = 32 // NSCodec always decodes to BGRA (4 bytes/pixel) + case 3: // RemoteFX (MS-RDPRFX) + if DecodeRemoteFX != nil { + pixels = DecodeRemoteFX(bitmapData, int(width), int(height)) + outBpp = 32 + } else { + slog.Warn("RemoteFX surface codec not available", "codecID", codecID) + return nil, nil + } + default: + slog.Warn("Unsupported surface codec", "codecID", codecID) + return nil, nil // skip unsupported codecs + } + + if pixels == nil { + return nil, nil + } + + // Flip vertically for bottom-up codecs. NSCodec decodes top-down but the + // bitmap coordinate system expects bottom-up. RFX (codecID=3) is already + // in the correct top-down orientation and must NOT be flipped. + if codecID != 3 { + stride := int(width) * int(outBpp) / 8 + h := int(height) + if stride > 0 && len(pixels) < stride*h { + // 解码产物不完整时翻转必然越界 panic;丢弃该帧而非崩溃。 + slog.Warn("surface bits: short pixel buffer, drop frame", + "codecID", codecID, "len", len(pixels), "want", stride*h) + return nil, nil + } + for y := 0; y < h/2; y++ { + top := y * stride + bot := (h - 1 - y) * stride + for i := range stride { + pixels[top+i], pixels[bot+i] = pixels[bot+i], pixels[top+i] + } + } + } + + return &BitmapData{ + DestLeft: destLeft, + DestTop: destTop, + DestRight: destRight, + DestBottom: destBottom, + Width: width, + Height: height, + BitsPerPixel: outBpp, + Flags: BITMAP_NO_PROCESSING, + BitmapLength: 0, + BitmapDataStream: pixels, + }, nil +} + +// decodeNSCodec decodes NSCodec (MS-RDPNSC) encoded bitmap data into BGRA pixels. +// Implements the decoder exactly as FreeRDP does (libfreerdp/codec/nsc.c). +func decodeNSCodec(data []byte, width, height int) []byte { + if len(data) < 20 { + slog.Warn("NSCodec data too short", "len", len(data)) + return nil + } + + r := bytes.NewReader(data) + lumaLen, _ := core.ReadUInt32LE(r) + orangeLen, _ := core.ReadUInt32LE(r) + greenLen, _ := core.ReadUInt32LE(r) + alphaLen, _ := core.ReadUInt32LE(r) + colorLossLevel, _ := core.ReadUInt8(r) + chromaSubsamplingLevel, _ := core.ReadUInt8(r) + _, _ = core.ReadUint16LE(r) // reserved + + if colorLossLevel < 1 { + colorLossLevel = 1 + } + shift := colorLossLevel - 1 + + slog.Debug("NSCodec", + "lumaLen", lumaLen, "orangeLen", orangeLen, + "greenLen", greenLen, "alphaLen", alphaLen, + "colorLossLevel", colorLossLevel, + "chromaSub", chromaSubsamplingLevel) + + remaining := data[20:] + + // Bounds check + totalPlaneLen := int(lumaLen + orangeLen + greenLen + alphaLen) + if totalPlaneLen > len(remaining) { + slog.Warn("NSCodec plane lengths exceed data", + "planeLens", totalPlaneLen, "available", len(remaining)) + return nil + } + + // Compute plane original (decompressed) sizes, matching FreeRDP: + // Y and A: tempWidth * height (Y uses rounded width for row stride) + // Co and Cg: (tempWidth>>1) * (tempHeight>>1) when chroma subsampled + tempWidth := (width + 7) &^ 7 // ROUND_UP_TO(width, 8) + tempHeight := (height + 1) &^ 1 // ROUND_UP_TO(height, 2) + + var yOrigSize, coOrigSize, cgOrigSize, aOrigSize int + if chromaSubsamplingLevel > 0 { + yOrigSize = tempWidth * height + coOrigSize = (tempWidth >> 1) * (tempHeight >> 1) + cgOrigSize = coOrigSize + } else { + yOrigSize = width * height + coOrigSize = yOrigSize + cgOrigSize = yOrigSize + } + aOrigSize = width * height + + // Acquire pooled plane buffers; released before returning so the pool is + // reused across decode calls without escaping to the caller. + yPlane := acquireNSCPlaneBuf(yOrigSize) + defer releaseNSCPlaneBuf(yPlane) + coPlane := acquireNSCPlaneBuf(coOrigSize) + defer releaseNSCPlaneBuf(coPlane) + cgPlane := acquireNSCPlaneBuf(cgOrigSize) + defer releaseNSCPlaneBuf(cgPlane) + + // Decompress each plane: if planeSize < originalSize → NRLE decode, + // if planeSize == 0 → fill with 0xFF, otherwise raw copy. + nscDecompressPlaneInto(remaining[:lumaLen], int(lumaLen), yPlane) + remaining = remaining[lumaLen:] + nscDecompressPlaneInto(remaining[:orangeLen], int(orangeLen), coPlane) + remaining = remaining[orangeLen:] + nscDecompressPlaneInto(remaining[:greenLen], int(greenLen), cgPlane) + remaining = remaining[greenLen:] + + var aPlane []byte + if alphaLen > 0 { + aPlane = acquireNSCPlaneBuf(aOrigSize) + defer releaseNSCPlaneBuf(aPlane) + nscDecompressPlaneInto(remaining[:alphaLen], int(alphaLen), aPlane) + } + + // YCoCg to BGRA conversion (matches FreeRDP nsc_decode exactly). + // FreeRDP formula: + // co_val = (INT16)(INT8)(((INT16)*coplane) << shift) + // cg_val = (INT16)(INT8)(((INT16)*cgplane) << shift) + // R = Y + co - cg + // G = Y + cg + // B = Y - co - cg + totalPixels := width * height + pixels := make([]byte, totalPixels*4) + + // Row widths for plane indexing (FreeRDP uses rw for Y, rw>>1 for chroma) + yRowWidth := width + coRowWidth := width + if chromaSubsamplingLevel > 0 { + yRowWidth = tempWidth + coRowWidth = tempWidth >> 1 + } + + if chromaSubsamplingLevel == 0 && aPlane == nil { + // Fast path: no chroma subsampling, no alpha override. + // ycoCgToBGRANoSub has SIMD implementations on amd64/arm64. + ycoCgToBGRANoSub(pixels, yPlane, coPlane, cgPlane, width*height, shift) + return pixels + } + + if chromaSubsamplingLevel > 0 { + // 2:1 horizontal chroma subsampling: each Co/Cg sample covers 2 pixels. + // Process 2 pixels per iteration to eliminate the px%2 modulo. + for py := range height { + yRowOff := py * yRowWidth + coIdx := (py >> 1) * coRowWidth + cgIdx := coIdx + outBase := py * width + + px := 0 + for ; px+1 < width; px += 2 { + coVal, cgVal := int16(0), int16(0) + if coIdx < len(coPlane) { + coVal = int16(int8(byte(int16(coPlane[coIdx]) << shift))) + } + if cgIdx < len(cgPlane) { + cgVal = int16(int8(byte(int16(cgPlane[cgIdx]) << shift))) + } + coIdx++ + cgIdx++ + + // Pixel px + off0 := (outBase + px) * 4 + yVal := int16(0) + if yIdx := yRowOff + px; yIdx < len(yPlane) { + yVal = int16(yPlane[yIdx]) + } + pixels[off0] = clampByte(yVal - coVal - cgVal) + pixels[off0+1] = clampByte(yVal + cgVal) + pixels[off0+2] = clampByte(yVal + coVal - cgVal) + if aPlane != nil && outBase+px < len(aPlane) { + pixels[off0+3] = aPlane[outBase+px] + } else { + pixels[off0+3] = 0xFF + } + + // Pixel px+1 (shares same Co/Cg sample) + off1 := off0 + 4 + yVal = int16(0) + if yIdx := yRowOff + px + 1; yIdx < len(yPlane) { + yVal = int16(yPlane[yIdx]) + } + pixels[off1] = clampByte(yVal - coVal - cgVal) + pixels[off1+1] = clampByte(yVal + cgVal) + pixels[off1+2] = clampByte(yVal + coVal - cgVal) + if aPlane != nil && outBase+px+1 < len(aPlane) { + pixels[off1+3] = aPlane[outBase+px+1] + } else { + pixels[off1+3] = 0xFF + } + } + // Handle odd width remainder + if px < width { + off := (outBase + px) * 4 + coVal, cgVal := int16(0), int16(0) + if coIdx < len(coPlane) { + coVal = int16(int8(byte(int16(coPlane[coIdx]) << shift))) + } + if cgIdx < len(cgPlane) { + cgVal = int16(int8(byte(int16(cgPlane[cgIdx]) << shift))) + } + yVal := int16(0) + if yIdx := yRowOff + px; yIdx < len(yPlane) { + yVal = int16(yPlane[yIdx]) + } + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + if aPlane != nil && outBase+px < len(aPlane) { + pixels[off+3] = aPlane[outBase+px] + } else { + pixels[off+3] = 0xFF + } + } + } + } else { + // No subsampling, but with alpha plane. + for py := range height { + yRowOff := py * yRowWidth + coIdx := py * coRowWidth + cgIdx := coIdx + outBase := py * width + + for px := range width { + yVal, coVal, cgVal := int16(0), int16(0), int16(0) + if yIdx := yRowOff + px; yIdx < len(yPlane) { + yVal = int16(yPlane[yIdx]) + } + if coIdx < len(coPlane) { + coVal = int16(int8(byte(int16(coPlane[coIdx]) << shift))) + } + if cgIdx < len(cgPlane) { + cgVal = int16(int8(byte(int16(cgPlane[cgIdx]) << shift))) + } + coIdx++ + cgIdx++ + + off := (outBase + px) * 4 + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + if outBase+px < len(aPlane) { + pixels[off+3] = aPlane[outBase+px] + } else { + pixels[off+3] = 0xFF + } + } + } + } + + return pixels +} + +func clampByte(v int16) uint8 { + if v < 0 { + return 0 + } + if v > 255 { + return 255 + } + return uint8(v) +} + +// nscDecompressPlane decompresses a single NSCodec plane. +// If planeSize == 0, fills with 0xFF. If planeSize >= originalSize, raw copy. +// Otherwise, uses the NRLE format (matching FreeRDP's nsc_rle_decode). +func nscDecompressPlane(input []byte, planeSize, originalSize int) []byte { + out := make([]byte, originalSize) + nscDecompressPlaneInto(input, planeSize, out) + return out +} + +// nscDecompressPlaneInto is the zero-allocation variant of nscDecompressPlane. +// out must be pre-allocated to exactly originalSize bytes. +func nscDecompressPlaneInto(input []byte, planeSize int, out []byte) { + originalSize := len(out) + if planeSize == 0 { + for i := range out { + out[i] = 0xFF + } + return + } + if planeSize >= originalSize { + copy(out, input[:originalSize]) + return + } + nrleDecodeInto(input[:planeSize], out) +} + +// nrleDecode decompresses NRLE (NSCodec Run-Length Encoding) data. +// Matches FreeRDP's nsc_rle_decode exactly: +// - 2 consecutive equal bytes trigger a run +// - If 3rd byte < 0xFF: run length = byte + 2 +// - If 3rd byte == 0xFF: run length = next 4 bytes as uint32 LE +// - Last 4 bytes of output are copied raw from input +func nrleDecode(input []byte, originalSize int) []byte { + output := make([]byte, originalSize) + nrleDecodeInto(input, output) + return output +} + +// nrleDecodeInto is the zero-allocation variant of nrleDecode. +// output must be pre-allocated to exactly originalSize bytes. +func nrleDecodeInto(input []byte, output []byte) { + originalSize := len(output) + left := originalSize + inPos := 0 + outPos := 0 + + for left > 4 && inPos < len(input) { + value := input[inPos] + inPos++ + + if left == 5 { + output[outPos] = value + outPos++ + left-- + } else if inPos < len(input) && value == input[inPos] { + // Run detected + inPos++ // skip the second occurrence + runLen := 0 + if inPos < len(input) { + if input[inPos] < 0xFF { + runLen = int(input[inPos]) + 2 + inPos++ + } else { + // Long run: skip 0xFF marker, read uint32 LE + inPos++ + if inPos+4 <= len(input) { + runLen = int(input[inPos]) | + int(input[inPos+1])<<8 | + int(input[inPos+2])<<16 | + int(input[inPos+3])<<24 + inPos += 4 + } + } + } + if runLen > left { + runLen = left + } + // Exponential-doubling copy for large runs is O(log n) instead of O(n). + n := min(runLen, originalSize-outPos) + output[outPos] = value + wrote := 1 + for wrote < n { + step := wrote + if wrote+step > n { + step = n - wrote + } + copy(output[outPos+wrote:outPos+wrote+step], output[outPos:outPos+wrote]) + wrote += step + } + outPos += n + left -= runLen + } else { + // Single byte + output[outPos] = value + outPos++ + left-- + } + } + + // Copy last 4 bytes raw + if left >= 4 && inPos+4 <= len(input) { + copy(output[outPos:outPos+4], input[inPos:inPos+4]) + } +} + +// TS_POINTER_NEW(慢路径 PTR_MSG_TYPE_POINTER 0x0008,FreeRDP +// update_read_pointer_new)与快速路径 FASTPATH_UPDATETYPE_POINTER 0x0B +// 共用同一布局:首个字段是 2 字节 xorBpp,其后才是 cacheIndex 等颜色 +// 指针属性(FreeRDP s_update_read_pointer_color:全部 u16 字段)。 +// TS_POINTER_NEW(慢路径 PTR_MSG_TYPE_POINTER 0x0008 与快速路径 +// FASTPATH_UPDATETYPE_POINTER 0x0B 共用布局,MS-RDPBCGR 2.2.9.1.1.4.4/4.5、 +// FreeRDP update_read_pointer_new):xorBpp 首字段,其后为颜色指针属性。 +// 可变长 blob 顺序按规范:xorMaskData(lengthXorMask 字节)在前, +// andMaskData(lengthAndMask 字节)在后——即 Data 先消费、Mask 后消费。 +// 注意 struc 按字段声明顺序消费,故 Data 必须声明在 Mask 之前; +// 此前 Mask 在前导致两个掩码整体互换,图像旋转错位产生花屏。 +type FastPathUpdatePointerPDU struct { + XorBpp uint16 `struc:"little"` + CacheIdx uint16 `struc:"little"` + HotX uint16 `struc:"little"` + HotY uint16 `struc:"little"` + Width uint16 `struc:"little"` + Height uint16 `struc:"little"` + MaskLen uint16 `struc:"little,sizeof=Mask"` // lengthAndMask + XorLen uint16 `struc:"little,sizeof=Data"` // lengthXorMask + Data []byte // xorMaskData(XOR 在前) + Mask []byte // andMaskData(AND 在后) +} + +func (*FastPathUpdatePointerPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_POINTER +} + +func (f *FastPathUpdatePointerPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, f) +} + +type FastPathPointerPositionPDU struct { + X uint16 `struc:"little"` + Y uint16 `struc:"little"` +} + +func (*FastPathPointerPositionPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_PTR_POSITION +} + +func (f *FastPathPointerPositionPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, f) +} + +type FastPathUpdatePointerNullPDU struct { +} + +func (*FastPathUpdatePointerNullPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_PTR_NULL +} +func (f *FastPathUpdatePointerNullPDU) Unpack(r io.Reader) error { + return nil +} + +// FastPathUpdatePointerDefaultPDU:FASTPATH_UPDATETYPE_PTR_DEFAULT(0x6), +// 无载荷——服务器要求客户端恢复默认系统箭头。此前该类型没有对应的 +// PDU 结构,分发时落入 PTR_POSITION 解析而报错丢弃,指针无法从 +// 隐藏/自定义形状切回箭头。 +type FastPathUpdatePointerDefaultPDU struct { +} + +func (*FastPathUpdatePointerDefaultPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_PTR_DEFAULT +} +func (f *FastPathUpdatePointerDefaultPDU) Unpack(r io.Reader) error { + return nil +} + +type FastPathUpdateCachedPDU struct { + CacheIdx uint16 `struc:"little"` +} + +func (*FastPathUpdateCachedPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_CACHED +} + +func (f *FastPathUpdateCachedPDU) Unpack(r io.Reader) error { + return struc.Unpack(r, f) +} + +type FastPathUpdatePDU struct { + UpdateHeader uint8 + Fragmentation uint8 + CompressionFlags uint8 + Size uint16 + Data UpdateData +} + +const ( + FASTPATH_OUTPUT_COMPRESSION_USED = 0x2 +) + +// Fast-path fragmentation bits (MS-RDPBCGR 2.2.9.1.1.3.1): SINGLE=0, +// FIRST=1, NEXT=2, LAST=3, left-shifted into bits 4-5 of the update header. +// 续片的 updateCode 无意义,重组后必须用首片记住的 code 解析。 +const ( + // FASTPATH_FRAGMENT_*:updateHeader 位 5-4 的分片字段(MS-RDPBCGR + // 2.2.9.1.1.3.1、FreeRDP fastpath.h 枚举左移 4 位后的线路值)。 + // SINGLE=0x0、LAST=0x1、FIRST=0x2、NEXT=0x3。此前 FIRST/NEXT/LAST + // 三个值轮换错位,导致首片被当孤儿丢弃、续片提前冲刷、尾片被缓存 + // 到下一条更新——SURFCMDS 流从命令中间开始,花屏与 resync 的总根源。 + FASTPATH_FRAGMENT_SINGLE = (0x0 << 4) + FASTPATH_FRAGMENT_LAST = (0x1 << 4) + FASTPATH_FRAGMENT_FIRST = (0x2 << 4) + FASTPATH_FRAGMENT_NEXT = (0x3 << 4) +) + +func readFastPathUpdatePDU(r io.Reader, code uint8) (*FastPathUpdatePDU, error) { + f := &FastPathUpdatePDU{} + var err error + var d UpdateData + switch code { + case FASTPATH_UPDATETYPE_ORDERS: + d = &FastPathOrdersPDU{} + case FASTPATH_UPDATETYPE_BITMAP: + d = &FastPathBitmapUpdateDataPDU{} + case FASTPATH_UPDATETYPE_PALETTE: + case FASTPATH_UPDATETYPE_SYNCHRONIZE: + case FASTPATH_UPDATETYPE_SURFCMDS: + //d = &FastPathSurfaceCmds{} + case FASTPATH_UPDATETYPE_PTR_NULL: + d = &FastPathUpdatePointerNullPDU{} + case FASTPATH_UPDATETYPE_PTR_DEFAULT: + d = &FastPathUpdatePointerDefaultPDU{} + case FASTPATH_UPDATETYPE_PTR_POSITION: + d = &FastPathPointerPositionPDU{} + case FASTPATH_UPDATETYPE_COLOR: + //d = &FastPathColorPdu{} + case FASTPATH_UPDATETYPE_CACHED: + d = &FastPathUpdateCachedPDU{} + case FASTPATH_UPDATETYPE_POINTER: + d = &FastPathUpdatePointerPDU{} + case FASTPATH_UPDATETYPE_LARGE_POINTER: + default: + return f, fmt.Errorf("Unknown FastPathPDU type 0x%x", code) + } + if d != nil { + err = d.Unpack(r) + if err != nil { + //slog.Error("Unpack:", err) + return nil, err + } + } else { + return nil, fmt.Errorf("Unsupport FastPathPDU type 0x%x", code) + } + + f.Data = d + return f, nil +} + +type ShareControlHeader struct { + TotalLength uint16 `struc:"little"` + PDUType uint16 `struc:"little"` + PDUSource uint16 `struc:"little"` +} + +type PDU struct { + ShareCtrlHeader *ShareControlHeader + Message PDUMessage +} + +func NewPDU(userId uint16, message PDUMessage) *PDU { + pdu := &PDU{} + pdu.ShareCtrlHeader = &ShareControlHeader{ + TotalLength: uint16(len(message.Serialize()) + 6), + PDUType: message.Type(), + PDUSource: userId, + } + pdu.Message = message + return pdu +} + +func readPDU(r io.Reader, mppc *core.MppcDecompressor) (*PDU, error) { + pdu := &PDU{} + var err error + header := &ShareControlHeader{} + err = struc.Unpack(r, header) + if err != nil { + return nil, err + } + + pdu.ShareCtrlHeader = header + + var d PDUMessage + switch pdu.ShareCtrlHeader.PDUType { + case PDUTYPE_DEMANDACTIVEPDU: + slog.Debug("readPDU:PDUTYPE_DEMANDACTIVEPDU") + d, err = readDemandActivePDU(r) + case PDUTYPE_DATAPDU: + slog.Debug("readPDU:PDUTYPE_DATAPDU") + d, err = readDataPDU(r, mppc) + case PDUTYPE_CONFIRMACTIVEPDU: + slog.Debug("readPDU:PDUTYPE_CONFIRMACTIVEPDU") + d, err = readConfirmActivePDU(r) + case PDUTYPE_DEACTIVATEALLPDU: + slog.Debug("readPDU:PDUTYPE_DEACTIVATEALLPDU") + d, err = readDeactiveAllPDU(r) + case PDUTYPE_SERVER_REDIR_PKT: + slog.Debug("readPDU:PDUTYPE_SERVER_REDIR_PKT") + d, err = readServerRedirectionPDU(r) + default: + slog.Error("PDU invalid pdu type", "type", fmt.Sprintf("0x%02x", pdu.ShareCtrlHeader.PDUType)) + } + if err != nil { + return nil, err + } + pdu.Message = d + return pdu, err +} + +func (p *PDU) serialize() []byte { + buff := &bytes.Buffer{} + struc.Pack(buff, p.ShareCtrlHeader) + core.WriteBytes(p.Message.Serialize(), buff) + return buff.Bytes() +} + +type SlowPathInputEvent struct { + EventTime uint32 `struc:"little"` + MessageType uint16 `struc:"little"` + Size int `struc:"skip"` + SlowPathInputData []byte `struc:"sizefrom=Size"` +} + +type PointerEvent struct { + PointerFlags uint16 `struc:"little"` + XPos uint16 `struc:"little"` + YPos uint16 `struc:"little"` +} + +func (p *PointerEvent) Serialize() []byte { + return []byte{ + byte(p.PointerFlags), byte(p.PointerFlags >> 8), + byte(p.XPos), byte(p.XPos >> 8), + byte(p.YPos), byte(p.YPos >> 8), + } +} + +// FastPathEncode appends this mouse event in the Fast-Path Input wire format +// (MS-RDPBCGR §2.2.8.1.2.2.3) to buf and returns the new slice. +func (p *PointerEvent) FastPathEncode(buf []byte) []byte { + buf = append(buf, byte(FASTPATH_INPUT_EVENT_MOUSE<<5)) + buf = append(buf, + byte(p.PointerFlags), byte(p.PointerFlags>>8), + byte(p.XPos), byte(p.XPos>>8), + byte(p.YPos), byte(p.YPos>>8)) + return buf +} + +type SynchronizeEvent struct { + Pad2Octets uint16 `struc:"little"` + ToggleFlags uint32 `struc:"little"` +} + +func (p *SynchronizeEvent) Serialize() []byte { + return []byte{ + byte(p.Pad2Octets), byte(p.Pad2Octets >> 8), + byte(p.ToggleFlags), byte(p.ToggleFlags >> 8), + byte(p.ToggleFlags >> 16), byte(p.ToggleFlags >> 24), + } +} + +type ScancodeKeyEvent struct { + KeyboardFlags uint16 `struc:"little"` + KeyCode uint16 `struc:"little"` + Pad2Octets uint16 `struc:"little"` +} + +func (p *ScancodeKeyEvent) Serialize() []byte { + return []byte{ + byte(p.KeyboardFlags), byte(p.KeyboardFlags >> 8), + byte(p.KeyCode), byte(p.KeyCode >> 8), + byte(p.Pad2Octets), byte(p.Pad2Octets >> 8), + } +} + +// FastPathEncode appends this scancode event in the Fast-Path Input wire +// format (MS-RDPBCGR §2.2.8.1.2.2.1) to buf and returns the new slice. +// +// Slow-path callers in this codebase historically encoded extended keys by +// stuffing the 0xE0 prefix into the high byte of KeyCode (e.g. 0xE048 for +// the up-arrow) and leaving KBDFLAGS_EXTENDED unset. Fast-path can only +// carry an 8-bit make code, so we promote any 0xE0XX encoding to the proper +// EXTENDED flag here. +func (p *ScancodeKeyEvent) FastPathEncode(buf []byte) []byte { + flags := byte(0) + if p.KeyboardFlags&KBDFLAGS_RELEASE != 0 { + flags |= FASTPATH_INPUT_KBDFLAGS_RELEASE + } + if p.KeyboardFlags&KBDFLAGS_EXTENDED != 0 || p.KeyCode&0xFF00 == 0xE000 { + flags |= FASTPATH_INPUT_KBDFLAGS_EXTENDED + } + if p.KeyboardFlags&KBDFLAGS_EXTENDED1 != 0 { + flags |= FASTPATH_INPUT_KBDFLAGS_EXTENDED1 + } + buf = append(buf, byte(FASTPATH_INPUT_EVENT_SCANCODE<<5)|flags) + buf = append(buf, byte(p.KeyCode)) + return buf +} + +type UnicodeKeyEvent struct { + KeyboardFlags uint16 `struc:"little"` + Unicode uint16 `struc:"little"` + Pad2Octets uint16 `struc:"little"` +} + +func (p *UnicodeKeyEvent) Serialize() []byte { + return []byte{ + byte(p.KeyboardFlags), byte(p.KeyboardFlags >> 8), + byte(p.Unicode), byte(p.Unicode >> 8), + byte(p.Pad2Octets), byte(p.Pad2Octets >> 8), + } +} + +// FastPathEncode appends this unicode key event in the Fast-Path Input wire +// format (MS-RDPBCGR §2.2.8.1.2.2.5) to buf and returns the new slice. +func (p *UnicodeKeyEvent) FastPathEncode(buf []byte) []byte { + flags := byte(0) + if p.KeyboardFlags&KBDFLAGS_RELEASE != 0 { + flags |= FASTPATH_INPUT_KBDFLAGS_RELEASE + } + buf = append(buf, byte(FASTPATH_INPUT_EVENT_UNICODE<<5)|flags) + buf = append(buf, byte(p.Unicode), byte(p.Unicode>>8)) + return buf +} + +type ClientInputEventPDU struct { + NumEvents uint16 `struc:"little,sizeof=SlowPathInputEvents"` + Pad2Octets uint16 `struc:"little"` + SlowPathInputEvents []SlowPathInputEvent `struc:"little"` +} + +func (*ClientInputEventPDU) Type2() uint8 { + return PDUTYPE2_INPUT +} +func (*ClientInputEventPDU) Unpack(io.Reader) error { + return nil +} + +// findSurfaceResync 从 from 开始扫描下一个合法的 SET/STREAM_SURFACE_BITS +// 命令头(cmdType 匹配 + width/height 与 dest 矩形一致 + bitmapDataLength +// 不越界,22/21 两种布局任一通过即认)。找不到返回 -1。 +func findSurfaceResync(buf []byte, from int) int { + for i := from; i+22 <= len(buf); i++ { + if !validSurfaceHeader(buf, i) { + continue + } + // 链式验证:候选头部之后必须紧跟缓冲结束、帧标记或另一个合法 + // 命令头。随机像素数据偶发形成单个伪头部可以,但连续两环全合 + // 法的概率实际为零——以此排除误命中画出的花屏瓦片。 + if total, ok := surfaceCmdTotal(buf, i); ok { + next := i + total + if next+22 > len(buf) || validSurfaceHeader(buf, next) || + isFrameMarkerAt(buf, next) || validSurfaceHeader(buf, next+8) { + return i + } + } + } + return -1 +} + +// isFrameMarkerAt 判断 pos 处是否为帧标记命令(cmdType=4,action 0/1)。 +func isFrameMarkerAt(buf []byte, pos int) bool { + return pos+8 <= len(buf) && + binary.LittleEndian.Uint16(buf[pos:]) == CMDTYPE_FRAME_MARKER && + binary.LittleEndian.Uint16(buf[pos+2:]) <= 1 +} + +// validSurfaceHeader 判断 pos 处是否为一个自洽的 SET/STREAM_SURFACE_BITS +// 命令头(cmdType 匹配 + codecID 已通告 + width/height 与 dest 一致 + +// bitmapDataLength 不越界,22/21 两种布局任一通过)。 +func validSurfaceHeader(buf []byte, pos int) bool { + if pos+22 > len(buf) { + return false + } + ct := binary.LittleEndian.Uint16(buf[pos:]) + if ct != CMDTYPE_SET_SURFACE_BITS && ct != CMDTYPE_STREAM_SURFACE_BITS { + return false + } + dl := int(binary.LittleEndian.Uint16(buf[pos+2:])) + dt := int(binary.LittleEndian.Uint16(buf[pos+4:])) + dr := int(binary.LittleEndian.Uint16(buf[pos+6:])) + db := int(binary.LittleEndian.Uint16(buf[pos+8:])) + remain := len(buf) - pos + fit := func(w, h int) bool { + return (w == dr-dl || w == dr-dl+1) && (h == db-dt || h == db-dt+1) + } + codec := func(b byte) bool { return b <= 3 } // 通告过的编解码族(0/1/3) + w22 := int(binary.LittleEndian.Uint16(buf[pos+14:])) + h22 := int(binary.LittleEndian.Uint16(buf[pos+16:])) + l22 := int(binary.LittleEndian.Uint32(buf[pos+18:])) + if l22 >= 0 && l22 <= 64<<20 && 22+l22 <= remain && fit(w22, h22) && codec(buf[pos+13]) { + return true + } + w21 := int(binary.LittleEndian.Uint16(buf[pos+13:])) + h21 := int(binary.LittleEndian.Uint16(buf[pos+15:])) + l21 := int(binary.LittleEndian.Uint32(buf[pos+17:])) + if l21 >= 0 && l21 <= 64<<20 && 21+l21 <= remain && fit(w21, h21) && codec(buf[pos+12]) { + return true + } + return false +} + +// surfaceCmdTotal 返回 pos 处合法命令的总长度(字节)。 +func surfaceCmdTotal(buf []byte, pos int) (int, bool) { + if pos+22 > len(buf) { + return 0, false + } + w22 := int(binary.LittleEndian.Uint16(buf[pos+14:])) + h22 := int(binary.LittleEndian.Uint16(buf[pos+16:])) + l22 := int(binary.LittleEndian.Uint32(buf[pos+18:])) + dl := int(binary.LittleEndian.Uint16(buf[pos+2:])) + dt := int(binary.LittleEndian.Uint16(buf[pos+4:])) + dr := int(binary.LittleEndian.Uint16(buf[pos+6:])) + db := int(binary.LittleEndian.Uint16(buf[pos+8:])) + fit := func(w, h int) bool { + return (w == dr-dl || w == dr-dl+1) && (h == db-dt || h == db-dt+1) + } + if l22 >= 0 && l22 <= 64<<20 && 22+l22 <= len(buf)-pos && fit(w22, h22) { + return 22 + l22, true + } + w21 := int(binary.LittleEndian.Uint16(buf[pos+13:])) + h21 := int(binary.LittleEndian.Uint16(buf[pos+15:])) + l21 := int(binary.LittleEndian.Uint32(buf[pos+17:])) + if l21 >= 0 && l21 <= 64<<20 && 21+l21 <= len(buf)-pos && fit(w21, h21) { + return 21 + l21, true + } + return 0, false +} + +// parseSurfaceCommandsIncremental 从缓冲解析完整的 surface 命令(可多条)。 +// 返回 consumed=完整消费的字节数;needMore=true 表示尾部命令不完整、 +// 需继续累积(此时 valid 恒为 true,缓冲必须保留);valid=false 表示 +// 缓冲开头不是合法命令流(调用方应整体丢弃)。dropped=true 表示尾部 +// 携带无法解析的未知结构(实测 Win10 整屏重绘批次末尾有 11 字节未文档 +// 化尾巴),已成功解析的前缀命令必须照常上屏、仅尾巴丢弃。 +// sticky:已锁定的 SET_SURFACE_BITS 头部长度(21/22;0=自动判定)。 +// 服务器布局会话内恒定,首条命令判定后必须锁定——逐帧启发式在 +// inclusive/exclusive 双兼容下会偶发选错 ±1 字节,错位累积到批次末尾 +// 即"未知命令类型→整段重置"(表现为周期性花屏+断流)。返回的 chosen +// 为本批实际判定的头部长度(sticky=0 时供调用方锁定)。 +func parseSurfaceCommandsIncremental(buf []byte, sticky int) (result SurfaceCommandsResult, consumed int, needMore bool, valid bool, dropped bool, chosen int) { + pos := 0 + chosen = sticky + for pos < len(buf) { + if len(buf)-pos < 2 { + return result, pos, true, true, false, chosen + } + cmdType := binary.LittleEndian.Uint16(buf[pos:]) + switch cmdType { + case CMDTYPE_SET_SURFACE_BITS, CMDTYPE_STREAM_SURFACE_BITS: + // 头部布局(字段偏移均从命令起始计): + // Win10 实测 22 字节:codecID@13、width@14、height@16、len u32@18、数据@22 + // FreeRDP 经典 21 字节:codecID@12、width@13、height@15、len u32@17、数据@21 + // 自动判定仅用于首条命令;锁定后按已知布局硬性校验。 + saneLen := func(l int) bool { return l >= 0 && l <= 64<<20 } + fitDim := func(w, h, dl, dt, dr, db int) bool { + // 兼容 inclusive/exclusive 两种 destRight/destBottom 约定 + return (w == dr-dl || w == dr-dl+1) && (h == db-dt || h == db-dt+1) + } + remain := len(buf) - pos + var hdrLen, rawLen int + switch sticky { + case 22: + if remain < 22 { + return result, pos, true, true, false, chosen + } + l := int(binary.LittleEndian.Uint32(buf[pos+18:])) + if !saneLen(l) { + return result, 0, false, false, false, chosen + } + if 22+l > remain { + return result, pos, true, true, false, chosen + } + hdrLen, rawLen = 22, l + case 21: + if remain < 21 { + return result, pos, true, true, false, chosen + } + l := int(binary.LittleEndian.Uint32(buf[pos+17:])) + if !saneLen(l) { + return result, 0, false, false, false, chosen + } + if 21+l > remain { + return result, pos, true, true, false, chosen + } + hdrLen, rawLen = 21, l + default: + if remain < 22 { + return result, pos, true, true, false, chosen + } + dl := int(binary.LittleEndian.Uint16(buf[pos+2:])) + dt := int(binary.LittleEndian.Uint16(buf[pos+4:])) + dr := int(binary.LittleEndian.Uint16(buf[pos+6:])) + db := int(binary.LittleEndian.Uint16(buf[pos+8:])) + w22 := int(binary.LittleEndian.Uint16(buf[pos+14:])) + h22 := int(binary.LittleEndian.Uint16(buf[pos+16:])) + l22 := int(binary.LittleEndian.Uint32(buf[pos+18:])) + w21 := int(binary.LittleEndian.Uint16(buf[pos+13:])) + h21 := int(binary.LittleEndian.Uint16(buf[pos+15:])) + l21 := int(binary.LittleEndian.Uint32(buf[pos+17:])) + switch { + case saneLen(l22) && 22+l22 <= remain && fitDim(w22, h22, dl, dt, dr, db): + hdrLen, rawLen, chosen = 22, l22, 22 + case saneLen(l21) && 21+l21 <= remain && fitDim(w21, h21, dl, dt, dr, db): + hdrLen, rawLen, chosen = 21, l21, 21 + case saneLen(l22) && 22+l22 <= remain: + hdrLen, rawLen, chosen = 22, l22, 22 + case saneLen(l21) && 21+l21 <= remain: + hdrLen, rawLen, chosen = 21, l21, 21 + case saneLen(l22) || saneLen(l21): + return result, pos, true, true, false, chosen // 数据未到齐,继续累积 + default: + return result, 0, false, false, false, chosen + } + } + total := hdrLen + rawLen + rect, err := decodeSurfaceBitsCmd(bytes.NewReader(buf[pos+2:pos+total]), hdrLen == 22) + if err != nil { + slog.Warn("decodeSurfaceBitsCmd", "err", err) + return result, 0, false, false, false, chosen + } + if rect != nil { + result.Rects = append(result.Rects, *rect) + } + pos += total + case CMDTYPE_FRAME_MARKER: + if len(buf)-pos < 2+2+4 { + return result, pos, true, true, false, chosen + } + if binary.LittleEndian.Uint16(buf[pos+2:]) == SURFCMD_FRAMEACTION_END { + result.FrameIDs = append(result.FrameIDs, binary.LittleEndian.Uint32(buf[pos+4:])) + } + pos += 8 + default: + // 重同步:未知命令类型说明流已错位(MPPC 历史或分片边界残留 + // 问题)。向前扫描下一个能通过 dest 一致性+长度双重校验的 + // SET/STREAM 命令头,只丢弃其之前的字节,避免整段缓冲报废。 + if next := findSurfaceResync(buf, pos+1); next > pos { + slog.Warn("surface cmd resync", "drop", next-pos, + "pos", pos, "next", next, "len", len(buf), + "at", fmt.Sprintf("% X", buf[pos:min(pos+24, len(buf))]), + "nextAt", fmt.Sprintf("% X", buf[next:min(next+24, len(buf))])) + pos = next + continue + } + slog.Warn("surface cmd unknown type", + "cmdType", fmt.Sprintf("0x%04X", cmdType), + "pos", pos, "len", len(buf), + "prefix", fmt.Sprintf("% X", buf[pos:min(pos+16, len(buf))])) + // 尾部未知结构(实测 Win10 整屏重绘批次末尾带 11 字节未文档 + // 化尾巴):已成功解析的前缀命令照常上屏,仅尾巴丢弃。 + // 若整批报废,每秒一次的悬浮部件更新会把整屏重绘全部丢掉 + //(表现为周期性花屏+断流)。 + return result, pos, false, true, true, chosen + } + } + return result, pos, false, true, false, chosen +} + +// PersistentKeyListPDU 是 TS_BITMAPCACHE_PERSISTENT_LIST_PDU +// (MS-RDPBCGR 2.2.2.3):连接 finalize 阶段声明客户端持久位图缓存 +// 持有的条目键,服务器重连后可据此免重传这些位图。 +// 头部固定 20 字节(numEntriesCache×5、totalEntriesCache×5、 +// bBitMask、pad1、pad3),随后为 numKeys×8 字节的键(key1 低 32 位、 +// key2 高 32 位,即 uint64 小端)。 +// +// 该 PDU 的键条目是变长且长度由 numEntriesCacheX 的分配决定,struc +// 反射无法直接表达,故手写序列化(与 PDUMessage 接口对接)。 +type PersistentKeyListPDU struct { + shareId uint32 + numEntries [5]uint16 + keys []uint64 +} + +func NewPersistentKeyListPDU(shareId uint32, keys []uint64, cellCounts [5]uint16) *PersistentKeyListPDU { + m := &PersistentKeyListPDU{shareId: shareId, keys: keys} + // 键按各缓存单元的广告容量依次分配(cache0 用完才进 cache1……), + // totalEntriesCacheX 回填为本 PDU 实际条数(FreeRDP 同款语义)。 + rest := uint16(len(keys)) + for i := 0; i < 5; i++ { + if rest < cellCounts[i] { + m.numEntries[i] = rest + } else { + m.numEntries[i] = cellCounts[i] + } + rest -= m.numEntries[i] + } + return m +} + +func (*PersistentKeyListPDU) Type() uint16 { + return PDUTYPE_DATAPDU +} + +func (m *PersistentKeyListPDU) Serialize() []byte { + payload := make([]byte, 24+len(m.keys)*8) + for i := 0; i < 5; i++ { + binary.LittleEndian.PutUint16(payload[i*2:], m.numEntries[i]) + // totalEntriesCacheX 写为实际条数(与 numEntries 一致) + binary.LittleEndian.PutUint16(payload[10+i*2:], m.numEntries[i]) + } + payload[20] = 0x03 // bBitMask = PERSIST_FIRST_PDU | PERSIST_LAST_PDU + // payload[21] pad1、payload[22:24] pad3 保持 0 + off := 24 + for _, k := range m.keys { + binary.LittleEndian.PutUint64(payload[off:], k) + off += 8 + } + hdr := NewShareDataHeader(len(payload), PDUTYPE2_BITMAPCACHE_PERSISTENT_LIST, m.shareId) + out := &bytes.Buffer{} + struc.Pack(out, hdr) + out.Write(payload) + return out.Bytes() +} diff --git a/protocol/pdu/data_test.go b/protocol/pdu/data_test.go new file mode 100644 index 0000000..7820394 --- /dev/null +++ b/protocol/pdu/data_test.go @@ -0,0 +1,327 @@ +package pdu + +import ( + "bytes" + "encoding/binary" + "testing" +) + +// buildSurfCmd 构造一条 SET/STREAM_SURFACE_BITS 命令。 +// win10=true 用 Win10 实测 22 字节头(codecID@13),否则用 +// FreeRDP 经典 21 字节头(codecID@12)。 +func buildSurfCmd(cmdType uint16, codecID, bpp byte, w, h uint16, payload []byte, win10 bool) []byte { + hdrLen := 21 + extra := 0 + if win10 { + hdrLen = 22 + extra = 1 + } + buf := make([]byte, hdrLen+len(payload)) + binary.LittleEndian.PutUint16(buf[0:], cmdType) + binary.LittleEndian.PutUint16(buf[2:], 10) // destLeft + binary.LittleEndian.PutUint16(buf[4:], 20) // destTop + binary.LittleEndian.PutUint16(buf[6:], 10+w-1) // destRight + binary.LittleEndian.PutUint16(buf[8:], 20+h-1) // destBottom + buf[10] = bpp // bpp + buf[11] = 0 // reserved + buf[12+extra] = codecID + binary.LittleEndian.PutUint16(buf[13+extra:], w) + binary.LittleEndian.PutUint16(buf[15+extra:], h) + binary.LittleEndian.PutUint32(buf[17+extra:], uint32(len(payload))) + copy(buf[hdrLen:], payload) + return buf +} + +func buildFrameMarker(action uint16, frameID uint32) []byte { + buf := make([]byte, 8) + binary.LittleEndian.PutUint16(buf[0:], CMDTYPE_FRAME_MARKER) + binary.LittleEndian.PutUint16(buf[2:], action) + binary.LittleEndian.PutUint32(buf[4:], frameID) + return buf +} + +// Win10 实测 22 字节头布局回归:cmdType(2)+dest(8)+bpp(1)+?@11+?@12+ +// codecID@13+width@14+height@16+len u32@18。偏移若错位, +// codecID/width/length 断言即失败。 +func TestParseSurfaceCommandsWin10Layout(t *testing.T) { + payload := make([]byte, 2*2*4) // 2x2 32bpp 原始像素 + buf := buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 0, 32, 2, 2, payload, true) + + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(buf, 0) + if !valid || needMore { + t.Fatalf("valid=%v needMore=%v", valid, needMore) + } + if consumed != len(buf) { + t.Fatalf("consumed=%d, want %d", consumed, len(buf)) + } + if len(result.Rects) != 1 { + t.Fatalf("rects=%d, want 1", len(result.Rects)) + } + r := result.Rects[0] + if r.Width != 2 || r.Height != 2 || r.BitsPerPixel != 32 { + t.Fatalf("size=%dx%d bpp=%d, want 2x2/32", r.Width, r.Height, r.BitsPerPixel) + } + if r.DestLeft != 10 || r.DestTop != 20 { + t.Fatalf("dest=(%d,%d), want (10,20)", r.DestLeft, r.DestTop) + } +} + +// FreeRDP 经典 21 字节头(codecID@12)的自适应回退。 +func TestParseSurfaceCommandsClassicLayout(t *testing.T) { + buf := buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 0, 32, 2, 2, + make([]byte, 2*2*4), false) + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(buf, 0) + if !valid || needMore || consumed != len(buf) { + t.Fatalf("valid=%v needMore=%v consumed=%d/%d", valid, needMore, consumed, len(buf)) + } + if len(result.Rects) != 1 { + t.Fatalf("rects=%d, want 1", len(result.Rects)) + } + if result.Rects[0].Width != 2 || result.Rects[0].Height != 2 { + t.Fatalf("size=%dx%d, want 2x2", result.Rects[0].Width, result.Rects[0].Height) + } +} + +// 真实抓包(Win10 19041, 24bpp 会话, 64x64 NSCodec 瓦片): +// 平面长度 4096+448+339+0=4883, bitmapDataLength=4903=20+4883。 +func TestParseSurfaceCommandsRealCapture(t *testing.T) { + // 真实命令前 48 字节(SURFCMD_DUMP 捕获),payload 截取自转储 + head := []byte{ + 0x01, 0x00, 0x40, 0x04, 0x40, 0x02, 0x80, 0x04, 0x80, 0x02, + 0x20, 0x00, 0x00, 0x01, 0x40, 0x00, 0x40, 0x00, 0x27, 0x13, + 0x00, 0x00, + } + nsHdr := []byte{ // NSCodec 流头:四平面长度+颜色参数 + 0x00, 0x10, 0x00, 0x00, 0xC0, 0x01, 0x00, 0x00, + 0x53, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x03, 0x01, 0x00, 0x00, + } + // bitmapData = NSCodec 流头(20) + 四平面数据(4883) = 4903 + buf := append(head, nsHdr...) + buf = append(buf, make([]byte, 4883)...) + if len(buf) != 22+4903 { + t.Fatalf("capture length %d, want %d", len(buf), 22+4903) + } + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(buf, 0) + if !valid || needMore { + t.Fatalf("valid=%v needMore=%v", valid, needMore) + } + if consumed != len(buf) { + t.Fatalf("consumed=%d, want %d", consumed, len(buf)) + } + if len(result.Rects) != 1 { + t.Fatalf("rects=%d, want 1 (codecID=1 NSCodec)", len(result.Rects)) + } + r := result.Rects[0] + if r.DestLeft != 1088 || r.DestTop != 576 || r.Width != 64 || r.Height != 64 { + t.Fatalf("rect=(%d,%d %dx%d), want (1088,576 64x64)", + r.DestLeft, r.DestTop, r.Width, r.Height) + } + if r.BitsPerPixel != 32 { + t.Fatalf("bpp=%d, want 32 (NSCodec 输出)", r.BitsPerPixel) + } +} + +func TestParseSurfaceCommandsRoundtrip(t *testing.T) { + payload := make([]byte, 2*2*4) // 2x2 32bpp 原始像素 + buf := append(buildFrameMarker(SURFCMD_FRAMEACTION_BEGIN, 7), + buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 0, 32, 2, 2, payload, true)...) + buf = append(buf, buildFrameMarker(SURFCMD_FRAMEACTION_END, 7)...) + + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(buf, 0) + if !valid { + t.Fatal("buffer should parse as valid command stream") + } + if needMore { + t.Fatal("complete stream should not need more data") + } + if consumed != len(buf) { + t.Fatalf("consumed=%d, want %d", consumed, len(buf)) + } + if len(result.Rects) != 1 { + t.Fatalf("rects=%d, want 1", len(result.Rects)) + } + r := result.Rects[0] + if r.DestLeft != 10 || r.DestTop != 20 { + t.Fatalf("dest=(%d,%d), want (10,20)", r.DestLeft, r.DestTop) + } + if r.Width != 2 || r.Height != 2 { + t.Fatalf("size=%dx%d, want 2x2", r.Width, r.Height) + } + if r.BitsPerPixel != 32 { + t.Fatalf("bpp=%d, want 32", r.BitsPerPixel) + } + if len(result.FrameIDs) != 1 || result.FrameIDs[0] != 7 { + t.Fatalf("frameIDs=%v, want [7]", result.FrameIDs) + } +} + +// Win10 会把一条命令拆到多个 fast-path PDU:前半段必须 needMore(且 +// 缓冲不得被调用方丢弃),后半段到齐后完整解析且 consumed 精确。 +func TestParseSurfaceCommandsSplitAcrossSegments(t *testing.T) { + cmd := buildSurfCmd(CMDTYPE_STREAM_SURFACE_BITS, 0, 32, 4, 2, + make([]byte, 4*2*4), true) + cut := 24 // 落在 payload 中间 + + r1, c1, more1, valid1, _, _ := parseSurfaceCommandsIncremental(cmd[:cut], 0) + if !valid1 || !more1 || c1 != 0 || len(r1.Rects) != 0 { + t.Fatalf("seg1: valid=%v more=%v consumed=%d rects=%d", + valid1, more1, c1, len(r1.Rects)) + } + + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(cmd, 0) + if !valid || needMore { + t.Fatalf("seg2: valid=%v needMore=%v", valid, needMore) + } + if consumed != len(cmd) || len(result.Rects) != 1 { + t.Fatalf("seg2: consumed=%d/%d rects=%d", + consumed, len(cmd), len(result.Rects)) + } +} + +// codecID=1 (NSCodec) 的 payload 即使解码失败,分帧也必须精确消费, +// 否则累积缓冲错位后整条流报废(黑屏的伴随症状)。 +func TestParseSurfaceCommandsNSCodecFraming(t *testing.T) { + cmd := buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 1, 32, 800, 480, + make([]byte, 100), true) // 非 NSCodec 合法数据 → 解码可能失败 + result, consumed, needMore, valid, _, _ := parseSurfaceCommandsIncremental(cmd, 0) + if !valid || needMore { + t.Fatalf("valid=%v needMore=%v", valid, needMore) + } + if consumed != len(cmd) { + t.Fatalf("consumed=%d, want %d", consumed, len(cmd)) + } + if len(result.Rects) > 1 { + t.Fatalf("rects=%d, want at most 1", len(result.Rects)) + } +} + +// 长度字段读到离谱值时应判无效并丢弃缓冲,而不是无限累积。 +func TestParseSurfaceCommandsRejectsBadLength(t *testing.T) { + cmd := buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 0, 32, 2, 2, + make([]byte, 16), true) + binary.LittleEndian.PutUint32(cmd[18:], 1<<30) // 1GB,非法 + _, _, _, valid, _, _ := parseSurfaceCommandsIncremental(cmd, 22) + if valid { + t.Fatal("absurd bitmapDataLength should invalidate the buffer") + } +} + +// TestFastPathFragmentConstants 锁定分片字段的线路值:MS-RDPBCGR +// 2.2.9.1.1.3.1 / FreeRDP fastpath.h 枚举(SINGLE=0, LAST=1, FIRST=2, +// NEXT=3)左移 4 位。此前 FIRST/NEXT/LAST 三个值轮换错位——首片被当 +// 孤儿丢弃、续片提前冲刷、尾片被缓存到下一条更新,是 24bpp 花屏与 +// surface cmd resync 的总根源。 +func TestFastPathFragmentConstants(t *testing.T) { + if FASTPATH_FRAGMENT_SINGLE != 0x00 { + t.Fatalf("SINGLE = %#x, want 0x00", FASTPATH_FRAGMENT_SINGLE) + } + if FASTPATH_FRAGMENT_LAST != 0x10 { + t.Fatalf("LAST = %#x, want 0x10", FASTPATH_FRAGMENT_LAST) + } + if FASTPATH_FRAGMENT_FIRST != 0x20 { + t.Fatalf("FIRST = %#x, want 0x20", FASTPATH_FRAGMENT_FIRST) + } + if FASTPATH_FRAGMENT_NEXT != 0x30 { + t.Fatalf("NEXT = %#x, want 0x30", FASTPATH_FRAGMENT_NEXT) + } +} + +// TestFastPathPointerNewLayout 用 Win10 实测指针字节锁定 TS_POINTER_NEW +// 布局:xorBpp u16 首字段 + 全 u16 颜色指针属性(FreeRDP +// update_read_pointer_new → s_update_read_pointer_color)。此前漏读 +// xorBpp 导致整体后移 2 字节,指针宽高/掩码长度全错位、光标无法显示。 +func TestFastPathPointerNewLayout(t *testing.T) { + raw := []byte{ + 0x20, 0x00, // xorBpp = 32 + 0x00, 0x00, // cacheIndex = 0 + 0x03, 0x00, // hotSpot.xPos = 3 + 0x03, 0x00, // hotSpot.yPos = 3 + 0x29, 0x00, // width = 41 + 0x27, 0x00, // height = 39 + 0xEA, 0x00, // lengthAndMask = 234 = ((41+15)/16)*2 * 39 + 0xFC, 0x18, // lengthXorMask = 6396 = ((41*32+15)/16)*2 * 39 + } + raw = append(raw, make([]byte, 234+6396)...) + p := &FastPathUpdatePointerPDU{} + if err := p.Unpack(bytes.NewReader(raw)); err != nil { + t.Fatalf("Unpack: %v", err) + } + if p.XorBpp != 32 || p.CacheIdx != 0 || p.HotX != 3 || p.HotY != 3 || + p.Width != 41 || p.Height != 39 { + t.Fatalf("header fields wrong: %+v", p) + } + if p.MaskLen != 234 || len(p.Mask) != 234 { + t.Fatalf("AND mask: len=%d MaskLen=%d, want 234", len(p.Mask), p.MaskLen) + } + if p.XorLen != 6396 || len(p.Data) != 6396 { + t.Fatalf("XOR mask: len=%d XorLen=%d, want 6396", len(p.Data), p.XorLen) + } +} + +// TestParseSurfaceCommandsDropsUnknownTail 验证尾部未知结构(实测 Win10 +// 整屏重绘批次末尾的 11 字节尾巴)不再拖垮整批:前缀命令照常解析上屏。 +func TestParseSurfaceCommandsDropsUnknownTail(t *testing.T) { + cmd := buildSurfCmd(CMDTYPE_SET_SURFACE_BITS, 0, 32, 2, 2, + make([]byte, 16), true) + bad := []byte{0x80, 0x00, 0x05, 0x00, 0x01, 0x00, 0x36, 0x01, 0x00, 0x00, 0x00} + buf := append(append([]byte{}, cmd...), bad...) + result, consumed, needMore, valid, dropped, _ := parseSurfaceCommandsIncremental(buf, 22) + if !valid { + t.Fatal("合法前缀不应判无效") + } + if !dropped { + t.Fatal("未知尾巴应标记 dropped") + } + if needMore { + t.Fatal("丢弃尾巴后不应再等待更多数据") + } + if consumed != len(cmd) { + t.Fatalf("consumed=%d,期望 %d", consumed, len(cmd)) + } + if len(result.Rects) != 1 { + t.Fatalf("前缀命令应产出 1 个矩形,实得 %d", len(result.Rects)) + } +} + +// TestPersistentKeyListPDU:键按单元容量分配、头 20 字节布局、 +// bBitMask=FIRST|LAST、键条目 uint64 小端,超 2042 键在发送层截断 +// (此处仅验证 PDU 自身,截断在 maybeSendPersistentKeyList)。 +func TestPersistentKeyListPDU(t *testing.T) { + // 2000 键、单元容量 600/1024/4096/0/0 → 600+1024+376=2000 + keys := make([]uint64, 2000) + for i := range keys { + keys[i] = uint64(i) + 0x10000 + } + cells := [5]uint16{600, 1024, 4096, 0, 0} + m := NewPersistentKeyListPDU(0x103EA, keys, cells) + if m.numEntries != [5]uint16{600, 1024, 376, 0, 0} { + t.Fatalf("numEntries=%v", m.numEntries) + } + b := m.Serialize() + // ShareDataHeader(12) + 头 24(num/total×10+bBitMask+pad)+ 2000×8 + if len(b) != 12+24+2000*8 { + t.Fatalf("len=%d", len(b)) + } + if b[12+20] != 0x03 || b[12+21] != 0 || b[12+22] != 0 { + t.Fatalf("bBitMask/pad wrong: %v", b[12+20:12+23]) + } + if got := binary.LittleEndian.Uint16(b[12:]); got != 600 { + t.Fatalf("numEntriesCache0=%d", got) + } + if got := binary.LittleEndian.Uint16(b[12+10:]); got != 600 { + t.Fatalf("totalEntriesCache0=%d", got) + } + if got := binary.LittleEndian.Uint16(b[12+4:]); got != 376 { + t.Fatalf("numEntriesCache2=%d", got) + } + // 第一条键 + if got := binary.LittleEndian.Uint64(b[12+24:]); got != 0x10000 { + t.Fatalf("first key=%x", got) + } + // 空键列表也合法(totalEntries 全 0) + empty := NewPersistentKeyListPDU(0x103EA, nil, cells).Serialize() + if len(empty) != 12+24 { + t.Fatalf("empty len=%d", len(empty)) + } +} diff --git a/protocol/pdu/orders.go b/protocol/pdu/orders.go new file mode 100644 index 0000000..546cc15 --- /dev/null +++ b/protocol/pdu/orders.go @@ -0,0 +1,1243 @@ +package pdu + +import ( + "bytes" + "errors" + "fmt" + "io" + "log/slog" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +type ControlFlag uint8 + +const ( + TS_STANDARD = 0x01 + TS_SECONDARY = 0x02 + TS_BOUNDS = 0x04 + TS_TYPE_CHANGE = 0x08 + TS_DELTA_COORDINATES = 0x10 + TS_ZERO_BOUNDS_DELTAS = 0x20 + TS_ZERO_FIELD_BYTE_BIT0 = 0x40 + TS_ZERO_FIELD_BYTE_BIT1 = 0x80 +) + +type PrimaryOrderType uint8 + +const ( + ORDER_TYPE_DSTBLT = 0x00 //0 + ORDER_TYPE_PATBLT = 0x01 //1 + ORDER_TYPE_SCRBLT = 0x02 //2 + //ORDER_TYPE_DRAWNINEGRID = 0x07 //7 + //ORDER_TYPE_MULTI_DRAWNINEGRID = 0x08 //8 + ORDER_TYPE_LINETO = 0x09 //9 + ORDER_TYPE_OPAQUERECT = 0x0A //10 + ORDER_TYPE_SAVEBITMAP = 0x0B //11 + ORDER_TYPE_MEMBLT = 0x0D //13 + ORDER_TYPE_MEM3BLT = 0x0E //14 + //ORDER_TYPE_MULTIDSTBLT = 0x0F //15 + //ORDER_TYPE_MULTIPATBLT = 0x10 //16 + //ORDER_TYPE_MULTISCRBLT = 0x11 //17 + //ORDER_TYPE_MULTIOPAQUERECT = 0x12 //18 + //ORDER_TYPE_FAST_INDEX = 0x13 //19 + ORDER_TYPE_POLYGON_SC = 0x14 //20 + ORDER_TYPE_POLYGON_CB = 0x15 //21 + ORDER_TYPE_POLYLINE = 0x16 //22 + //ORDER_TYPE_FAST_GLYPH = 0x18 //24 + ORDER_TYPE_ELLIPSE_SC = 0x19 //25 + ORDER_TYPE_ELLIPSE_CB = 0x1A //26 + ORDER_TYPE_TEXT2 = 0x1B //27 +) + +type SecondaryOrderType uint8 + +const ( + ORDER_TYPE_BITMAP_UNCOMPRESSED = 0x00 + ORDER_TYPE_CACHE_COLOR_TABLE = 0x01 + ORDER_TYPE_CACHE_BITMAP_COMPRESSED = 0x02 + ORDER_TYPE_CACHE_GLYPH = 0x03 + ORDER_TYPE_BITMAP_UNCOMPRESSED_V2 = 0x04 + ORDER_TYPE_BITMAP_COMPRESSED_V2 = 0x05 + ORDER_TYPE_CACHE_BRUSH = 0x07 + ORDER_TYPE_BITMAP_COMPRESSED_V3 = 0x08 +) + +func (s SecondaryOrderType) String() string { + name := "Unknown" + switch s { + case ORDER_TYPE_BITMAP_UNCOMPRESSED: + name = "Cache Bitmap" + case ORDER_TYPE_CACHE_COLOR_TABLE: + name = "Cache Color Table" + case ORDER_TYPE_CACHE_BITMAP_COMPRESSED: + name = "Cache Bitmap (Compressed)" + case ORDER_TYPE_CACHE_GLYPH: + name = "Cache Glyph" + case ORDER_TYPE_BITMAP_UNCOMPRESSED_V2: + name = "Cache Bitmap V2" + case ORDER_TYPE_BITMAP_COMPRESSED_V2: + name = "Cache Bitmap V2 (Compressed)" + case ORDER_TYPE_CACHE_BRUSH: + name = "Cache Brush" + case ORDER_TYPE_BITMAP_COMPRESSED_V3: + name = "Cache Bitmap V3" + } + return fmt.Sprintf("[0x%02d] %s", s, name) +} + +/* Alternate Secondary Drawing Orders */ +const ( + ORDER_TYPE_SWITCH_SURFACE = 0x00 + ORDER_TYPE_CREATE_OFFSCREEN_BITMAP = 0x01 + ORDER_TYPE_STREAM_BITMAP_FIRST = 0x02 + ORDER_TYPE_STREAM_BITMAP_NEXT = 0x03 + ORDER_TYPE_CREATE_NINE_GRID_BITMAP = 0x04 + ORDER_TYPE_GDIPLUS_FIRST = 0x05 + ORDER_TYPE_GDIPLUS_NEXT = 0x06 + ORDER_TYPE_GDIPLUS_END = 0x07 + ORDER_TYPE_GDIPLUS_CACHE_FIRST = 0x08 + ORDER_TYPE_GDIPLUS_CACHE_NEXT = 0x09 + ORDER_TYPE_GDIPLUS_CACHE_END = 0x0A + ORDER_TYPE_WINDOW = 0x0B + ORDER_TYPE_COMPDESK_FIRST = 0x0C + ORDER_TYPE_FRAME_MARKER = 0x0D +) + +const ( + GLYPH_FRAGMENT_NOP = 0x00 + GLYPH_FRAGMENT_USE = 0xFE + GLYPH_FRAGMENT_ADD = 0xFF + + CBR2_HEIGHT_SAME_AS_WIDTH = 0x01 + CBR2_PERSISTENT_KEY_PRESENT = 0x02 + CBR2_NO_BITMAP_COMPRESSION_HDR = 0x08 + CBR2_DO_NOT_CACHE = 0x10 +) + +const ( + ORDER_PRIMARY = iota + ORDER_SECONDARY + ORDER_ALTSEC +) + +type OrderPdu struct { + ControlFlags uint8 + Type int + Altsec *Altsec + Primary *Primary + Secondary *Secondary + // CacheBitmapV2 非 nil 时携带该次级订单解析出的位图缓存条目 + //(stage6 6.4b:位图缓存 v2 的存储来源) + CacheBitmapV2 *CacheBitmapV2Order +} + +func (o *OrderPdu) HasBounds() bool { + return o.ControlFlags&TS_BOUNDS != 0 +} + +type Altsec struct { +} + +type Secondary struct { +} + +type Primary struct { + Bounds Bounds + Data PrimaryOrder +} + +type FastPathOrdersPDU struct { + NumberOrders uint16 + OrderPdus []OrderPdu +} + +func (*FastPathOrdersPDU) FastPathUpdateType() uint8 { + return FASTPATH_UPDATETYPE_ORDERS +} + +func (f *FastPathOrdersPDU) Unpack(r io.Reader) error { + f.NumberOrders, _ = core.ReadUint16LE(r) + //slog.Debug("NumberOrders:", f.NumberOrders) + for i := 0; i < int(f.NumberOrders); i++ { + var o OrderPdu + o.ControlFlags, _ = core.ReadUInt8(r) + if o.ControlFlags&TS_STANDARD == 0 { + //slog.Debug("Altsec order") + o.processAltsecOrder(r) + o.Type = ORDER_ALTSEC + //return errors.New("Not support") + } else if o.ControlFlags&TS_SECONDARY != 0 { + //slog.Debug("Secondary order") + o.processSecondaryOrder(r) + o.Type = ORDER_SECONDARY + } else { + //slog.Debug("Primary order") + o.processPrimaryOrder(r) + o.Type = ORDER_PRIMARY + } + + if f.OrderPdus == nil { + f.OrderPdus = make([]OrderPdu, 0, f.NumberOrders) + } + f.OrderPdus = append(f.OrderPdus, o) + } + return nil +} +func (o *OrderPdu) processAltsecOrder(r io.Reader) error { + orderType := o.ControlFlags >> 2 + //slog.Debug("Altsec:", orderType) + switch orderType { + case ORDER_TYPE_SWITCH_SURFACE: + case ORDER_TYPE_CREATE_OFFSCREEN_BITMAP: + case ORDER_TYPE_STREAM_BITMAP_FIRST: + case ORDER_TYPE_STREAM_BITMAP_NEXT: + case ORDER_TYPE_CREATE_NINE_GRID_BITMAP: + case ORDER_TYPE_GDIPLUS_FIRST: + case ORDER_TYPE_GDIPLUS_NEXT: + case ORDER_TYPE_GDIPLUS_END: + case ORDER_TYPE_GDIPLUS_CACHE_FIRST: + case ORDER_TYPE_GDIPLUS_CACHE_NEXT: + case ORDER_TYPE_GDIPLUS_CACHE_END: + case ORDER_TYPE_WINDOW: + case ORDER_TYPE_COMPDESK_FIRST: + case ORDER_TYPE_FRAME_MARKER: + core.ReadUInt32LE(r) + } + + return nil +} +func (o *OrderPdu) processSecondaryOrder(r io.Reader) error { + var sec Secondary + length, _ := core.ReadUint16LE(r) + flags, _ := core.ReadUint16LE(r) + orderType, _ := core.ReadUInt8(r) + + slog.Debug("processSecondaryOrder", "SecondaryOrderType", SecondaryOrderType(orderType)) + + b, _ := core.ReadBytes(int(length)+13-6, r) + r0 := bytes.NewReader(b) + + switch orderType { + case ORDER_TYPE_BITMAP_UNCOMPRESSED: + fallthrough + case ORDER_TYPE_CACHE_BITMAP_COMPRESSED: + compressed := (orderType == ORDER_TYPE_CACHE_BITMAP_COMPRESSED) + sec.updateCacheBitmapOrder(r0, compressed, flags) + case ORDER_TYPE_BITMAP_UNCOMPRESSED_V2: + fallthrough + case ORDER_TYPE_BITMAP_COMPRESSED_V2: + compressed := (orderType == ORDER_TYPE_BITMAP_COMPRESSED_V2) + if cbv2 := sec.updateCacheBitmapV2Order(r0, compressed, flags); cbv2 != nil { + // 位图缓存 v2(stage6 6.4b):交由上层存储/引用回贴 + o.CacheBitmapV2 = cbv2 + } + case ORDER_TYPE_BITMAP_COMPRESSED_V3: + sec.updateCacheBitmapV3Order(r0, flags) + case ORDER_TYPE_CACHE_COLOR_TABLE: + sec.updateCacheColorTableOrder(r0, flags) + case ORDER_TYPE_CACHE_GLYPH: + sec.updateCacheGlyphOrder(r0, flags) + case ORDER_TYPE_CACHE_BRUSH: + sec.updateCacheBrushOrder(r0, flags) + default: + slog.Debug("processSecondaryOrder", "Unsupport order type", orderType) + } + + return nil +} +func (b *Bounds) updateBounds(r io.Reader) { + present, _ := core.ReadUInt8(r) + + if present&1 != 0 { + readOrderCoord(r, &b.left, false) + } else if present&16 != 0 { + readOrderCoord(r, &b.left, true) + } + + if present&2 != 0 { + readOrderCoord(r, &b.top, false) + } else if present&32 != 0 { + readOrderCoord(r, &b.top, true) + } + + if present&4 != 0 { + readOrderCoord(r, &b.right, false) + } else if present&64 != 0 { + readOrderCoord(r, &b.right, true) + } + if present&8 != 0 { + readOrderCoord(r, &b.bottom, false) + } else if present&128 != 0 { + readOrderCoord(r, &b.bottom, true) + } +} + +type PrimaryOrder interface { + Type() int + Unpack(io.Reader, uint32, bool) error +} + +var ( + orderType uint8 + bounds Bounds +) + +func (o *OrderPdu) processPrimaryOrder(r io.Reader) error { + o.Primary = &Primary{} + if o.ControlFlags&TS_TYPE_CHANGE != 0 { + orderType, _ = core.ReadUInt8(r) + } + size := 1 + switch orderType { + case ORDER_TYPE_MEM3BLT, ORDER_TYPE_TEXT2: + size = 3 + + case ORDER_TYPE_PATBLT, ORDER_TYPE_MEMBLT, ORDER_TYPE_LINETO, ORDER_TYPE_POLYGON_CB, ORDER_TYPE_ELLIPSE_CB: + size = 2 + } + + if o.ControlFlags&TS_ZERO_FIELD_BYTE_BIT0 != 0 { + size-- + } + if o.ControlFlags&TS_ZERO_FIELD_BYTE_BIT1 != 0 { + if size < 2 { + size = 0 + } else { + size -= 2 + } + } + var present uint32 + for i := 0; i < size; i++ { + bits, _ := core.ReadUInt8(r) + present |= uint32(bits) << (i * 8) + } + + if o.ControlFlags&TS_BOUNDS != 0 { + if o.ControlFlags&TS_ZERO_BOUNDS_DELTAS == 0 { + bounds.updateBounds(r) + } + //slog.Debug("updateBounds") + o.Primary.Bounds = bounds + } + + delta := o.ControlFlags&TS_DELTA_COORDINATES != 0 + + //slog.Debug(fmt.Sprintf("present=%d,delta=%v", present, delta)) + + var p PrimaryOrder + switch orderType { + case ORDER_TYPE_DSTBLT: + p = &Dstblt{} + + case ORDER_TYPE_PATBLT: + p = &Patblt{} + + case ORDER_TYPE_SCRBLT: + p = &Scrblt{} + + //case ORDER_TYPE_DRAWNINEGRID: + + //case ORDER_TYPE_MULTI_DRAWNINEGRID: + + case ORDER_TYPE_LINETO: + p = &LineTo{} + + case ORDER_TYPE_OPAQUERECT: + p = &OpaqueRect{} + + case ORDER_TYPE_SAVEBITMAP: + p = &SaveBitmap{} + + case ORDER_TYPE_MEMBLT: + p = &Memblt{} + + case ORDER_TYPE_MEM3BLT: + p = &Mem3blt{} + + //case ORDER_TYPE_MULTIDSTBLT: + + //case ORDER_TYPE_MULTIPATBLT: + + //case ORDER_TYPE_MULTISCRBLT: + + //case ORDER_TYPE_MULTIOPAQUERECT: + + //case ORDER_TYPE_FAST_INDEX: + + case ORDER_TYPE_POLYGON_SC: + p = &PolygonSc{} + + case ORDER_TYPE_POLYGON_CB: + p = &PolygonCb{} + + case ORDER_TYPE_POLYLINE: + p = &Polyline{} + + //case ORDER_TYPE_FAST_GLYPH: + + case ORDER_TYPE_ELLIPSE_SC: + p = &EllipeSc{} + + case ORDER_TYPE_ELLIPSE_CB: + p = &EllipeCb{} + + case ORDER_TYPE_TEXT2: + p = &GlayphIndex{} + default: + slog.Error("processPrimaryOrder", "orderType", orderType) + return errors.New("Not Support order type") + } + if p != nil { + if err := p.Unpack(r, present, delta); err != nil { + return err + } + } + + o.Primary.Data = p + return nil +} +func readOrderCoord(r io.Reader, coord *int32, delta bool) { + if delta { + change, _ := core.ReadUInt8(r) + *coord += int32(int8(change)) + } else { + change, _ := core.ReadUint16LE(r) + *coord = int32(int16(change)) + } +} + +type Dstblt struct { + x int32 + y int32 + cx int32 + cy int32 + opcode uint8 +} + +func (d *Dstblt) Type() int { + return ORDER_TYPE_DSTBLT +} +func (d *Dstblt) Unpack(r io.Reader, present uint32, delta bool) error { + slog.Debug("Dstblt Order") + if present&0x01 != 0 { + readOrderCoord(r, &d.x, delta) + } + if present&0x02 != 0 { + readOrderCoord(r, &d.y, delta) + } + if present&0x04 != 0 { + readOrderCoord(r, &d.cx, delta) + } + if present&0x08 != 0 { + readOrderCoord(r, &d.cy, delta) + } + if present&0x10 != 0 { + d.opcode, _ = core.ReadUInt8(r) + } + return nil +} + +type Patblt struct { + x int32 + y int32 + cx int32 + cy int32 + opcode uint8 + bgcolour [4]uint8 + fgcolour [4]uint8 + brush Brush +} + +func (d *Patblt) Type() int { + return ORDER_TYPE_PATBLT +} +func (d *Patblt) Unpack(r io.Reader, present uint32, delta bool) error { + slog.Debug("Patblt Order") + if present&0x01 != 0 { + readOrderCoord(r, &d.x, delta) + } + if present&0x02 != 0 { + readOrderCoord(r, &d.y, delta) + } + if present&0x04 != 0 { + readOrderCoord(r, &d.cx, delta) + } + if present&0x08 != 0 { + readOrderCoord(r, &d.cy, delta) + } + if present&0x10 != 0 { + d.opcode, _ = core.ReadUInt8(r) + } + if present&0x0020 != 0 { + b, g, r, a := updateReadColorRef(r) + d.bgcolour[0], d.bgcolour[1], d.bgcolour[2], d.bgcolour[3] = b, g, r, a + } + if present&0x0040 != 0 { + b, g, r, a := updateReadColorRef(r) + d.fgcolour[0], d.fgcolour[1], d.fgcolour[2], d.fgcolour[3] = b, g, r, a + } + d.brush.updateBrush(r, present>>7) + + return nil +} + +type Brush struct { + X uint8 + Y uint8 + Style uint8 + Hatch uint8 + Data []byte +} + +func (b *Brush) updateBrush(r io.Reader, present uint32) { + if present&1 != 0 { + b.X, _ = core.ReadUInt8(r) + } + + if present&2 != 0 { + b.Y, _ = core.ReadUInt8(r) + } + + if present&4 != 0 { + b.Style, _ = core.ReadUInt8(r) + } + + if present&8 != 0 { + b.Hatch, _ = core.ReadUInt8(r) + } + + if present&16 != 0 { + data, _ := core.ReadBytes(7, r) + b.Data = make([]byte, 0, 8) + b.Data = append(b.Data, b.Hatch) + b.Data = append(b.Data, data...) + } +} + +type Scrblt struct { + X int32 + Y int32 + Cx int32 + Cy int32 + Opcode uint8 + Srcx int32 + Srcy int32 +} + +func (d *Scrblt) Type() int { + return ORDER_TYPE_SCRBLT +} + +var d Scrblt + +func (d1 *Scrblt) Unpack(r io.Reader, present uint32, delta bool) error { + slog.Debug("Scrblt Order") + if present&0x0001 != 0 { + readOrderCoord(r, &d.X, delta) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.Y, delta) + } + if present&0x0004 != 0 { + readOrderCoord(r, &d.Cx, delta) + } + if present&0x0008 != 0 { + readOrderCoord(r, &d.Cy, delta) + } + if present&0x0010 != 0 { + d.Opcode, _ = core.ReadUInt8(r) + } + if present&0x0020 != 0 { + readOrderCoord(r, &d.Srcx, delta) + } + if present&0x0040 != 0 { + readOrderCoord(r, &d.Srcy, delta) + } + *d1 = d + return nil +} + +type LineTo struct { + Mixmode uint16 + Startx int32 + Starty int32 + Endx int32 + Endy int32 + Bgcolour [4]uint8 + Opcode uint8 + Pen Pen +} + +func (d *LineTo) Type() int { + return ORDER_TYPE_LINETO +} +func (d *LineTo) Unpack(r io.Reader, present uint32, delta bool) error { + slog.Debug("LineTo Order") + if present&0x0001 != 0 { + d.Mixmode, _ = core.ReadUint16LE(r) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.Startx, delta) + } + if present&0x0004 != 0 { + readOrderCoord(r, &d.Starty, delta) + } + if present&0x008 != 0 { + readOrderCoord(r, &d.Endx, delta) + } + if present&0x0010 != 0 { + readOrderCoord(r, &d.Endy, delta) + } + if present&0x0020 != 0 { + b, g, r, a := updateReadColorRef(r) + d.Bgcolour[0], d.Bgcolour[1], d.Bgcolour[2], d.Bgcolour[3] = b, g, r, a + } + if present&0x0040 != 0 { + d.Opcode, _ = core.ReadUInt8(r) + } + + d.Pen.updatePen(r, present>>7) + + return nil +} + +type Pen struct { + Style uint8 + Width uint8 + Colour [4]uint8 +} + +func (d *Pen) updatePen(r io.Reader, present uint32) { + if present&1 != 0 { + d.Style, _ = core.ReadUInt8(r) + } + + if present&2 != 0 { + d.Width, _ = core.ReadUInt8(r) + } + + if present&4 != 0 { + b, g, r, a := updateReadColorRef(r) + d.Colour[0], d.Colour[1], d.Colour[2], d.Colour[3] = b, g, r, a + } +} + +type OpaqueRect struct { + X int32 + Y int32 + Cx int32 + Cy int32 + Colour [4]uint8 +} + +func (d *OpaqueRect) Type() int { + return ORDER_TYPE_OPAQUERECT +} +func (d *OpaqueRect) Unpack(r io.Reader, present uint32, delta bool) error { + slog.Debug("OpaqueRect Order") + if present&0x0001 != 0 { + readOrderCoord(r, &d.X, delta) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.Y, delta) + } + if present&0x0004 != 0 { + readOrderCoord(r, &d.Cx, delta) + } + if present&0x0008 != 0 { + readOrderCoord(r, &d.Cy, delta) + } + if present&0x0010 != 0 { + i, _ := core.ReadUInt8(r) + d.Colour[0] = i + } + if present&0x0020 != 0 { + i, _ := core.ReadUInt8(r) + d.Colour[1] = i + } + if present&0x0040 != 0 { + i, _ := core.ReadUInt8(r) + d.Colour[2] = i + } + return nil +} + +type SaveBitmap struct { + Offset uint32 + Left int32 + Top int32 + Right int32 + Bottom int32 + action uint8 +} + +func (d *SaveBitmap) Type() int { + return ORDER_TYPE_SAVEBITMAP +} +func (d *SaveBitmap) Unpack(r io.Reader, present uint32, delta bool) error { + if present&0x0001 != 0 { + d.Offset, _ = core.ReadUInt32LE(r) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.Left, delta) + } + if present&0x0004 != 0 { + readOrderCoord(r, &d.Top, delta) + } + if present&0x0008 != 0 { + readOrderCoord(r, &d.Right, delta) + } + if present&0x0010 != 0 { + readOrderCoord(r, &d.Bottom, delta) + } + if present&0x0020 != 0 { + d.action, _ = core.ReadUInt8(r) + } + return nil +} + +type Memblt struct { + ColourTable uint8 + CacheId uint8 + X int32 + Y int32 + Cx int32 + Cy int32 + Opcode uint8 + Srcx int32 + Srcy int32 + CacheIdx uint16 +} + +func (d *Memblt) Type() int { + return ORDER_TYPE_MEMBLT +} +func (d *Memblt) Unpack(r io.Reader, present uint32, delta bool) error { + if present&0x0001 != 0 { + d.CacheId, _ = core.ReadUInt8(r) + d.ColourTable, _ = core.ReadUInt8(r) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.X, delta) + } + if present&0x0004 != 0 { + readOrderCoord(r, &d.Y, delta) + } + if present&0x0008 != 0 { + readOrderCoord(r, &d.Cx, delta) + } + if present&0x0010 != 0 { + readOrderCoord(r, &d.Cy, delta) + } + if present&0x0020 != 0 { + d.Opcode, _ = core.ReadUInt8(r) + } + if present&0x0040 != 0 { + readOrderCoord(r, &d.Srcx, delta) + } + if present&0x0080 != 0 { + readOrderCoord(r, &d.Srcy, delta) + } + if present&0x0100 != 0 { + d.CacheIdx, _ = core.ReadUint16LE(r) + } + return nil +} + +type Mem3blt struct { + ColourTable uint8 + CacheId uint8 + X int32 + Y int32 + Cx int32 + Cy int32 + Opcode uint8 + Srcx int32 + Srcy int32 + Bgcolour [4]uint8 + Fgcolour [4]uint8 + Brush Brush + CacheIdx uint16 +} + +func (d *Mem3blt) Type() int { + return ORDER_TYPE_MEM3BLT +} +func (d *Mem3blt) Unpack(r io.Reader, present uint32, delta bool) error { + if present&0x000001 != 0 { + d.CacheId, _ = core.ReadUInt8(r) + d.ColourTable, _ = core.ReadUInt8(r) + } + if present&0x000002 != 0 { + readOrderCoord(r, &d.X, delta) + } + if present&0x000004 != 0 { + readOrderCoord(r, &d.Y, delta) + } + if present&0x000008 != 0 { + readOrderCoord(r, &d.Cx, delta) + } + if present&0x000010 != 0 { + readOrderCoord(r, &d.Cy, delta) + } + if present&0x000020 != 0 { + d.Opcode, _ = core.ReadUInt8(r) + } + if present&0x000040 != 0 { + readOrderCoord(r, &d.Srcx, delta) + } + if present&0x000080 != 0 { + readOrderCoord(r, &d.Srcy, delta) + } + if present&0x000100 != 0 { + b, g, r, a := updateReadColorRef(r) + d.Bgcolour[0], d.Bgcolour[1], d.Bgcolour[2], d.Bgcolour[3] = b, g, r, a + } + if present&0x000200 != 0 { + b, g, r, a := updateReadColorRef(r) + d.Fgcolour[0], d.Fgcolour[1], d.Fgcolour[2], d.Fgcolour[3] = b, g, r, a + } + d.Brush.updateBrush(r, present>>10) + if present&0x008000 != 0 { + d.CacheIdx, _ = core.ReadUint16LE(r) + } + if present&0x010000 != 0 { + core.ReadUint16LE(r) + } + + return nil +} + +type PolygonSc struct { + X int32 + Y int32 + Opcode uint8 + Fillmode uint8 + Fgcolour [4]uint8 + Npoints uint8 + Points []Point +} + +type Point struct { + X int32 + Y int32 +} + +func (d *PolygonSc) Type() int { + return ORDER_TYPE_POLYGON_SC +} +func (d *PolygonSc) Unpack(r io.Reader, present uint32, delta bool) error { + if present&0x0001 != 0 { + readOrderCoord(r, &d.X, delta) + } + if present&0x0002 != 0 { + readOrderCoord(r, &d.Y, delta) + } + if present&0x0004 != 0 { + d.Opcode, _ = core.ReadUInt8(r) + } + if present&0x0008 != 0 { + d.Fillmode, _ = core.ReadUInt8(r) + } + if present&0x0010 != 0 { + b, g, r, a := updateReadColorRef(r) + d.Fgcolour[0], d.Fgcolour[1], d.Fgcolour[2], d.Fgcolour[3] = b, g, r, a + } + if present&0x0020 != 0 { + d.Npoints, _ = core.ReadUInt8(r) + d.Points = make([]Point, 0, d.Npoints+1) + } + if present&0x0040 != 0 { + size, _ := core.ReadUInt8(r) + data, _ := core.ReadBytes(int(size), r) + d.Points = append(d.Points, Point{d.X, d.Y}) + var flags uint8 + r = bytes.NewReader(data) + for i := 1; i <= int(d.Npoints); i++ { + var p Point + if (i-1)%4 == 0 { + flags, _ = core.ReadUInt8(r) + } + if (^flags)&0x80 != 0 { + p.X = parseDelta(r) + } + if (^flags)&0x40 != 0 { + p.Y = parseDelta(r) + } + flags <<= 2 + } + } + + return nil +} + +func parseDelta(r io.Reader) (v int32) { + b, _ := core.ReadUInt8(r) + if b&0x40 != 0 { + v = int32(b) | (^0x3F) + } else { + v = int32(b & 0x3F) + } + if b&0x80 != 0 { + b, _ := core.ReadUInt8(r) + v = (v << 8) | int32(b) + } + return +} + +type PolygonCb struct { +} + +func (d *PolygonCb) Type() int { + return ORDER_TYPE_POLYGON_CB +} +func (d *PolygonCb) Unpack(r io.Reader, present uint32, delta bool) error { + return nil +} + +type Polyline struct { +} + +func (d *Polyline) Type() int { + return ORDER_TYPE_POLYLINE +} +func (d *Polyline) Unpack(r io.Reader, present uint32, delta bool) error { + return nil +} + +type EllipeSc struct { +} + +func (d *EllipeSc) Type() int { + return ORDER_TYPE_ELLIPSE_SC +} +func (d *EllipeSc) Unpack(r io.Reader, present uint32, delta bool) error { + return nil +} + +type EllipeCb struct { +} + +func (d *EllipeCb) Type() int { + return ORDER_TYPE_ELLIPSE_CB +} +func (d *EllipeCb) Unpack(r io.Reader, present uint32, delta bool) error { + return nil +} + +type GlayphIndex struct { +} + +func (d *GlayphIndex) Type() int { + return ORDER_TYPE_TEXT2 +} +func (d *GlayphIndex) Unpack(r io.Reader, present uint32, delta bool) error { + return nil +} + +/*Secondary*/ +func (s *Secondary) updateCacheBitmapOrder(r io.Reader, compressed bool, flags uint16) { + var cb CacheBitmapOrder + cb.cacheId, _ = core.ReadUInt8(r) + core.ReadUInt8(r) + cb.bitmapWidth, _ = core.ReadUInt8(r) + cb.bitmapHeight, _ = core.ReadUInt8(r) + cb.bitmapBpp, _ = core.ReadUInt8(r) + bitmapLength, _ := core.ReadUint16LE(r) + cb.cacheIndex, _ = core.ReadUint16LE(r) + var bitmapComprHdr []byte + if compressed { + if (flags & NO_BITMAP_COMPRESSION_HDR) == 0 { + bitmapComprHdr, _ = core.ReadBytes(8, r) + bitmapLength -= 8 + } + } + cb.bitmapComprHdr = bitmapComprHdr + cb.bitmapDataStream, _ = core.ReadBytes(int(bitmapLength), r) + cb.bitmapLength = bitmapLength + +} + +type CacheBitmapOrder struct { + cacheId uint8 + bitmapBpp uint8 + bitmapWidth uint8 + bitmapHeight uint8 + bitmapLength uint16 + cacheIndex uint16 + bitmapComprHdr []byte + bitmapDataStream []byte +} + +func getCbV2Bpp(bpp uint32) (b uint32) { + switch bpp { + case 3: + b = 8 + case 4: + b = 16 + case 5: + b = 24 + case 6: + b = 32 + default: + b = 0 + } + return +} + +type CacheBitmapV2Order struct { + CacheId uint32 + Flags uint32 + Key1 uint32 + Key2 uint32 + BitmapBpp uint32 + BitmapWidth uint8 + BitmapHeight uint8 + BitmapLength uint16 + CacheIndex uint32 + Compressed bool + CbCompFirstRowSize uint16 + CbCompMainBodySize uint16 + CbScanWidth uint16 + CbUncompressedSize uint16 + BitmapDataStream []byte +} + +func (s *Secondary) updateCacheBitmapV2Order(r io.Reader, compressed bool, flags uint16) *CacheBitmapV2Order { + var cb CacheBitmapV2Order + cb.CacheId = uint32(flags) & 0x0003 + cb.Flags = (uint32(flags) & 0xFF80) >> 7 + bitsPerPixelId := (uint32(flags) & 0x0078) >> 3 + cb.BitmapBpp = getCbV2Bpp(bitsPerPixelId) + + if cb.Flags&CBR2_PERSISTENT_KEY_PRESENT != 0 { + cb.Key1, _ = core.ReadUInt32LE(r) + cb.Key2, _ = core.ReadUInt32LE(r) + } + + if cb.Flags&CBR2_HEIGHT_SAME_AS_WIDTH != 0 { + cb.BitmapWidth, _ = core.ReadUInt8(r) + cb.BitmapHeight = cb.BitmapWidth + } else { + cb.BitmapWidth, _ = core.ReadUInt8(r) + cb.BitmapHeight, _ = core.ReadUInt8(r) + } + + bitmapLength, _ := core.ReadUint16LE(r) + cacheIndex, _ := core.ReadUInt8(r) + + if cb.Flags&CBR2_DO_NOT_CACHE != 0 { + cb.CacheIndex = 0x7FFF + } else { + cb.CacheIndex = uint32(cacheIndex) + } + + if compressed { + if cb.Flags&CBR2_NO_BITMAP_COMPRESSION_HDR == 0 { + cb.CbCompFirstRowSize, _ = core.ReadUint16LE(r) + cb.CbCompMainBodySize, _ = core.ReadUint16LE(r) + cb.CbScanWidth, _ = core.ReadUint16LE(r) + cb.CbUncompressedSize, _ = core.ReadUint16LE(r) + bitmapLength = cb.CbCompMainBodySize + } + } + + cb.BitmapDataStream, _ = core.ReadBytes(int(bitmapLength), r) + cb.BitmapLength = bitmapLength + cb.Compressed = compressed + + return &cb +} + +type CacheBitmapV3Order struct { + cacheId uint32 + bpp uint32 + flags uint32 + cacheIndex uint16 + key1 uint32 + key2 uint32 + bitmapData BitmapDataEx +} +type BitmapDataEx struct { + bpp uint8 + codecID uint8 + width uint16 + height uint16 + length uint32 + data []byte +} + +func (s *Secondary) updateCacheBitmapV3Order(r io.Reader, flags uint16) { + var cb CacheBitmapV3Order + + cb.cacheId = uint32(flags) & 0x00000003 + cb.flags = (uint32(flags) & 0x0000FF80) >> 7 + bitsPerPixelId := (uint32(flags) & 0x00000078) >> 3 + cb.bpp = getCbV2Bpp(bitsPerPixelId) + + cacheIndex, _ := core.ReadUint16LE(r) + cb.cacheIndex = cacheIndex + cb.key1, _ = core.ReadUInt32LE(r) + cb.key2, _ = core.ReadUInt32LE(r) + + bitmapData := &cb.bitmapData + bitmapData.bpp, _ = core.ReadUInt8(r) + core.ReadUInt8(r) + core.ReadUInt8(r) + bitmapData.codecID, _ = core.ReadUInt8(r) + bitmapData.width, _ = core.ReadUint16LE(r) + bitmapData.height, _ = core.ReadUint16LE(r) + new_len, _ := core.ReadUInt32LE(r) + + bitmapData.data, _ = core.ReadBytes(int(new_len), r) + bitmapData.length = new_len + +} + +type CacheColorTableOrder struct { + cacheIndex uint8 + numberColors uint16 + colorTable [256 * 4]uint8 +} + +func (s *Secondary) updateCacheColorTableOrder(r io.Reader, flags uint16) { + var cb CacheColorTableOrder + cb.cacheIndex, _ = core.ReadUInt8(r) + cb.numberColors, _ = core.ReadUint16LE(r) + + if cb.numberColors != 256 { + /* This field MUST be set to 256 */ + return + } + + for i := 0; i < int(cb.numberColors)*4; i++ { + cb.colorTable[i], cb.colorTable[i+1], cb.colorTable[i+2], cb.colorTable[i+3] = updateReadColorRef(r) + } +} +func updateReadColorRef(r io.Reader) (uint8, uint8, uint8, uint8) { + blue, _ := core.ReadUInt8(r) + green, _ := core.ReadUInt8(r) + red, _ := core.ReadUInt8(r) + core.ReadUInt8(r) + + return blue, green, red, 255 +} + +type CacheGlyphOrder struct { + cacheId uint8 + nglyphs uint8 + glyphs []CacheGlyph +} +type CacheGlyph struct { + character uint16 + offset uint16 + baseline uint16 + width uint16 + height uint16 + datasize int + data []uint8 +} + +func (s *Secondary) updateCacheGlyphOrder(r io.Reader, flags uint16) { + var cb CacheGlyphOrder + + cb.cacheId, _ = core.ReadUInt8(r) + cb.nglyphs, _ = core.ReadUInt8(r) + cb.glyphs = make([]CacheGlyph, 0, cb.nglyphs) + + for i := 0; i < int(cb.nglyphs); i++ { + var c CacheGlyph + c.character, _ = core.ReadUint16LE(r) + c.offset, _ = core.ReadUint16LE(r) + c.baseline, _ = core.ReadUint16LE(r) + c.width, _ = core.ReadUint16LE(r) + c.height, _ = core.ReadUint16LE(r) + + c.datasize = int(c.height*((c.width+7)/8)+3) & ^3 + c.data, _ = core.ReadBytes(c.datasize, r) + + cb.glyphs = append(cb.glyphs, c) + } +} + +type CacheBrushOrder struct { + index uint8 + bpp uint8 + cx uint8 + cy uint8 + style uint8 + length uint8 + data []uint8 +} + +func (s *Secondary) updateCacheBrushOrder(r io.Reader, flags uint16) { + var cb CacheBrushOrder + cb.index, _ = core.ReadUInt8(r) + cb.bpp, _ = core.ReadUInt8(r) + cb.cx, _ = core.ReadUInt8(r) + cb.cy, _ = core.ReadUInt8(r) + cb.style, _ = core.ReadUInt8(r) + cb.length, _ = core.ReadUInt8(r) + if cb.cx == 8 && cb.cy == 8 { + if cb.bpp == 1 { + for i := 7; i >= 0; i-- { + cb.data[i], _ = core.ReadUInt8(r) + } + } else { + bpp := int(cb.bpp) - 2 + if int(cb.length) == 16+4*bpp { + /* compressed brush */ + data, _ := core.ReadBytes(int(cb.length), r) + cb.data = update_decompress_brush(data, bpp) + } else { + /* uncompressed brush */ + scanline := 8 * 8 * bpp + cb.data, _ = core.ReadBytes(scanline, r) + } + } + } +} +func update_decompress_brush(in []uint8, bpp int) []uint8 { + var pal_index, in_index, shift int + + pal := in[16:] + out := make([]uint8, 8*8*bpp) + /* read it bottom up */ + for y := 7; y >= 0; y-- { + /* 2 bytes per row */ + x := 0 + for range 2 { + /* 4 pixels per byte */ + shift = 6 + for shift >= 0 { + pal_index = int((in[in_index] >> shift) & 3) + /* size of palette entries depends on bpp */ + for i := range bpp { + out[(y*8+x)*bpp+i] = pal[pal_index*bpp+i] + } + x++ + shift -= 2 + } + in_index++ + } + } + + return out +} + +/*Primary*/ +type Bounds struct { + left int32 + top int32 + right int32 + bottom int32 +} +type OrderInfo struct { + controlFlags uint32 + orderType uint32 + fieldFlags uint32 + boundsFlags uint32 + bounds Bounds + deltaCoordinates bool +} diff --git a/protocol/pdu/pdu.go b/protocol/pdu/pdu.go new file mode 100644 index 0000000..3e49e48 --- /dev/null +++ b/protocol/pdu/pdu.go @@ -0,0 +1,935 @@ +package pdu + +import ( + "bytes" + "encoding/binary" + "fmt" + "log/slog" + "sync" + "sync/atomic" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/gcc" +) + +var readerPool = sync.Pool{ + New: func() any { return new(bytes.Reader) }, +} + +// LastServerErrorInfo 记录服务器最近一次通过 SET_ERROR_INFO PDU 报告的错误码。 +// 服务器断开前通常先发送该 PDU(如 0x0000112F = ERRINFO_GRAPHICS_SUBSYSTEM_FAILED), +// 上层据此判断断连原因并做降级重试。 +var LastServerErrorInfo atomic.Uint32 + +// fastPathBufPool reuses byte slices for serializing fast-path input PDUs. +// Capacity 128 covers the maximum frame: 1 + 7*15 = 106 bytes. +var fastPathBufPool = sync.Pool{ + New: func() any { return make([]byte, 0, 128) }, +} + +type PDULayer struct { + bitmapCachePersistKeys []uint64 + emission.Emitter + transport core.Transport + sharedId uint32 + userId uint16 + channelId uint16 + serverCapabilities map[CapsType]Capability + clientCapabilities map[CapsType]Capability + fastPathSender core.FastPathSender + // serverFastPathInput is set after capability exchange when both sides + // advertise INPUT_FLAG_FASTPATH_INPUT, allowing client input to be sent + // using the much shorter fast-path framing (MS-RDPBCGR §2.2.8.1.2). + serverFastPathInput bool + demandActivePDU *DemandActivePDU + mppc *core.MppcDecompressor +} + +func NewPDULayer(t core.Transport) *PDULayer { + p := &PDULayer{ + Emitter: *emission.NewEmitter(), + transport: t, + sharedId: 0x103EA, + serverCapabilities: map[CapsType]Capability{ + CAPSTYPE_GENERAL: &GeneralCapability{ + ProtocolVersion: 0x0200, + }, + CAPSTYPE_BITMAP: &BitmapCapability{ + Receive1BitPerPixel: 0x0001, + Receive4BitsPerPixel: 0x0001, + Receive8BitsPerPixel: 0x0001, + BitmapCompressionFlag: 0x0001, + MultipleRectangleSupport: 0x0001, + }, + CAPSTYPE_ORDER: &OrderCapability{ + DesktopSaveXGranularity: 1, + DesktopSaveYGranularity: 20, + MaximumOrderLevel: 1, + OrderFlags: NEGOTIATEORDERSUPPORT, + DesktopSaveSize: 480 * 480, + }, + CAPSTYPE_POINTER: &PointerCapability{ColorPointerCacheSize: 20}, + CAPSTYPE_INPUT: &InputCapability{}, + CAPSTYPE_VIRTUALCHANNEL: &VirtualChannelCapability{}, + CAPSTYPE_FONT: &FontCapability{SupportFlags: 0x0001}, + CAPSTYPE_COLORCACHE: &ColorCacheCapability{CacheSize: 0x0006}, + CAPSTYPE_SHARE: &ShareCapability{}, + }, + clientCapabilities: map[CapsType]Capability{ + CAPSTYPE_GENERAL: &GeneralCapability{ + ProtocolVersion: 0x0200, + }, + CAPSTYPE_BITMAP: &BitmapCapability{ + Receive1BitPerPixel: 0x0001, + Receive4BitsPerPixel: 0x0001, + Receive8BitsPerPixel: 0x0001, + BitmapCompressionFlag: 0x0001, + MultipleRectangleSupport: 0x0001, + }, + CAPSTYPE_ORDER: &OrderCapability{ + DesktopSaveXGranularity: 1, + DesktopSaveYGranularity: 20, + MaximumOrderLevel: 1, + OrderFlags: NEGOTIATEORDERSUPPORT, + DesktopSaveSize: 480 * 480, + TextANSICodePage: 0x4e4, + }, + CAPSTYPE_CONTROL: &ControlCapability{0, 0, 2, 2}, + CAPSTYPE_ACTIVATION: &WindowActivationCapability{}, + CAPSTYPE_POINTER: &PointerCapability{1, 20, 20}, + CAPSTYPE_SHARE: &ShareCapability{}, + CAPSTYPE_COLORCACHE: &ColorCacheCapability{6, 0}, + CAPSTYPE_SOUND: &SoundCapability{0x0001, 0}, + CAPSTYPE_INPUT: &InputCapability{}, + CAPSTYPE_FONT: &FontCapability{0x0001, 0}, + CAPSTYPE_BRUSH: &BrushCapability{BRUSH_COLOR_8x8}, + CAPSTYPE_GLYPHCACHE: &GlyphCapability{}, + CAPSETTYPE_BITMAP_CODECS: newClientBitmapCodecsCapability(), + CAPSTYPE_BITMAPCACHE_REV2: &BitmapCache2Capability{ + BitmapCachePersist: 2, + CachesNum: 5, + BmpC0Cells: 0x258, + BmpC1Cells: 0x258, + BmpC2Cells: 0x800, + BmpC3Cells: 0x1000, + BmpC4Cells: 0x800, + }, + CAPSTYPE_VIRTUALCHANNEL: &VirtualChannelCapability{0, 1600}, + CAPSETTYPE_MULTIFRAGMENTUPDATE: &MultiFragmentUpdate{0x3F0000}, + CAPSTYPE_RAIL: &RemoteProgramsCapability{ + RailSupportLevel: RAIL_LEVEL_SUPPORTED | + RAIL_LEVEL_SHELL_INTEGRATION_SUPPORTED | + RAIL_LEVEL_LANGUAGE_IME_SYNC_SUPPORTED | + RAIL_LEVEL_SERVER_TO_CLIENT_IME_SYNC_SUPPORTED | + RAIL_LEVEL_HIDE_MINIMIZED_APPS_SUPPORTED | + RAIL_LEVEL_WINDOW_CLOAKING_SUPPORTED | + RAIL_LEVEL_HANDSHAKE_EX_SUPPORTED | + RAIL_LEVEL_DOCKED_LANGBAR_SUPPORTED, + }, + CAPSETTYPE_LARGE_POINTER: &LargePointerCapability{1}, + CAPSETTYPE_COMPDESK: &DesktopCompositionCapability{ + CompDeskSupportLevel: 1, // COMPDESK_SUPPORTED + }, + CAPSETTYPE_SURFACE_COMMANDS: &SurfaceCommandsCapability{ + CmdFlags: SURFCMDS_SET_SURFACE_BITS | SURFCMDS_STREAM_SURFACE_BITS | SURFCMDS_FRAME_MARKER, + }, + CAPSSETTYPE_FRAME_ACKNOWLEDGE: &FrameAcknowledgeCapability{2}, + }, + mppc: core.NewMppcDecompressor(), + } + + t.On("close", func() { + p.Emit("close") + }).On("error", func(err error) { + p.Emit("error", err) + }) + return p +} + +func (p *PDULayer) sendPDU(message PDUMessage) { + pdu := NewPDU(p.userId, message) + p.transport.Write(pdu.serialize()) +} + +func (p *PDULayer) sendDataPDU(message DataPDUData) { + dataPdu := NewDataPDU(message, p.sharedId) + p.sendPDU(dataPdu) +} + +func (p *PDULayer) SetFastPathSender(f core.FastPathSender) { + p.fastPathSender = f +} + +type Client struct { + *PDULayer + clientCoreData *gcc.ClientCoreData + buff *bytes.Buffer + // fragCode 记住分片更新首片的 updateCode:续片的 updateCode 无意义 + //(Win10 发送的是载荷字节),重组完成后必须用它来解析。 + fragCode uint8 + // fragActive/fragPoisoned:分片重组进行中标记;解压失败的首片会把 + // 整个更新标记为中毒,其剩余分片被吞掉(防止孤儿分片拼出垃圾)。 + fragActive bool + fragPoisoned bool + // decErrLogs:MPPC 解压失败日志限次计数器。 + decErrLogs int + // surfResetDumped:SURFCMDS 重置全量转储限次(离线分析用)。 + surfResetDumped int + // bmpRectLogs:16bpp 矩形取证日志限次计数器。 + bmpRectLogs int + // bitmapCachePersistKeys:客户端持久位图缓存键(6.4b M2), + // finalize 时经 maybeSendPersistentKeyList 上报服务器。 + bitmapCachePersistKeys []uint64 + // bmpFrame:fast-path 位图跨段累积缓冲(Win10 会把一个位图更新拆到 + // 多个 fast-path PDU,按矩形结构完整解析后整体消费)。 + bmpFrame []byte + // surfBuf:fast-path 表面命令跨段累积缓冲(按命令流自定界解析)。 + surfBuf []byte + // surfHdrLen:SET_SURFACE_BITS 头部长度(21=经典 / 22=Win10),由本 + // 连接第一条命令判定后锁定。服务器布局会话内恒定;逐帧启发式在 + // inclusive/exclusive 双兼容下会偶发选错 ±1 字节,错位累积到批次 + // 末尾即"未知命令类型→整段重置"——表现为周期性花屏+断流。 + surfHdrLen int +} + +func NewClient(t core.Transport) *Client { + c := &Client{ + PDULayer: NewPDULayer(t), + buff: &bytes.Buffer{}, + } + c.transport.Once("connect", c.connect) + return c +} + +func (c *Client) connect(data *gcc.ClientCoreData, userId uint16, channelId uint16) { + slog.Debug("pdu connect", "userId", userId, "channelId", channelId) + c.clientCoreData = data + c.userId = userId + c.channelId = channelId + c.transport.Once("data", c.recvDemandActivePDU) +} + +func (c *Client) recvDemandActivePDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + pdu, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvDemandActivePDU", "err", err) + return + } + if pdu.ShareCtrlHeader.PDUType != PDUTYPE_DEMANDACTIVEPDU { + if pdu.ShareCtrlHeader.PDUType == PDUTYPE_DEACTIVATEALLPDU { + // Per [MS-RDPBCGR] the server may send DeactivateAllPDU before + // DemandActivePDU (e.g. GNOME RDP after RDPGFX capability + // exchange). Stay on the same connection and keep waiting, + // exactly as FreeRDP does. + slog.Debug("received DeactivateAllPDU while waiting for DemandActivePDU; continuing to wait") + c.transport.Once("data", c.recvDemandActivePDU) + return + } + if pdu.ShareCtrlHeader.PDUType == PDUTYPE_SERVER_REDIR_PKT { + if redir, ok := pdu.Message.(*ServerRedirectionPDU); ok { + c.Emit("redirect", redir) + } + return + } + slog.Debug("ignore message during connection sequence", "type", pdu.ShareCtrlHeader.PDUType) + c.transport.Once("data", c.recvDemandActivePDU) + return + } + c.sharedId = pdu.Message.(*DemandActivePDU).SharedId + c.demandActivePDU = pdu.Message.(*DemandActivePDU) + for _, caps := range c.demandActivePDU.CapabilitySets { + slog.Debug("serverCaps", "type", caps.Type(), "value", caps) + c.serverCapabilities[caps.Type()] = caps + } + if ic, ok := c.serverCapabilities[CAPSTYPE_INPUT].(*InputCapability); ok { + c.serverFastPathInput = ic.Flags&INPUT_FLAG_FASTPATH_INPUT != 0 + } + + c.sendConfirmActivePDU() + c.sendClientFinalizeSynchronizePDU() + c.transport.Once("data", c.recvServerSynchronizePDU) +} + +func (c *Client) sendConfirmActivePDU() { + pdu := NewConfirmActivePDU() + generalCapa := c.clientCapabilities[CAPSTYPE_GENERAL].(*GeneralCapability) + generalCapa.OSMajorType = OSMAJORTYPE_WINDOWS + generalCapa.OSMinorType = OSMINORTYPE_WINDOWS_NT + generalCapa.GeneralCompressionTypes = 0x0002 // PACKET_COMPR_TYPE_64K: advertise MPPC-64K support + generalCapa.ExtraFlags = LONG_CREDENTIALS_SUPPORTED | NO_BITMAP_COMPRESSION_HDR | + FASTPATH_OUTPUT_SUPPORTED | AUTORECONNECT_SUPPORTED + generalCapa.RefreshRectSupport = 1 + generalCapa.SuppressOutputSupport = 1 + + bitmapCapa := c.clientCapabilities[CAPSTYPE_BITMAP].(*BitmapCapability) + bitmapCapa.PreferredBitsPerPixel = 32 + bitmapCapa.DesktopWidth = c.clientCoreData.DesktopWidth + bitmapCapa.DesktopHeight = c.clientCoreData.DesktopHeight + bitmapCapa.DesktopResizeFlag = 0x0001 + + orderCapa := c.clientCapabilities[CAPSTYPE_ORDER].(*OrderCapability) + orderCapa.OrderFlags = NEGOTIATEORDERSUPPORT | ZEROBOUNDSDELTASSUPPORT | COLORINDEXSUPPORT | ORDERFLAGS_EXTRA_FLAGS + orderCapa.OrderSupportExFlags |= ORDERFLAGS_EX_ALTSEC_FRAME_MARKER_SUPPORT + orderCapa.OrderSupport[TS_NEG_DSTBLT_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_PATBLT_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_SCRBLT_INDEX] = 1 + // MEMBLT + 位图缓存 v2(stage6 6.4b M1 会话内缓存):服务器把重复 + // 位图存入客户端缓存单元,后续以 MemBlt 引用回贴,替代全量重发。 + orderCapa.OrderSupport[TS_NEG_MEMBLT_INDEX] = 1 + //orderCapa.OrderSupport[TS_NEG_LINETO_INDEX] = 1 + //orderCapa.OrderSupport[TS_NEG_MEM3BLT_INDEX] = 1 + //orderCapa.OrderSupport[TS_NEG_POLYLINE_INDEX] = 1 + /*orderCapa.OrderSupport[TS_NEG_MULTIOPAQUERECT_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_GLYPH_INDEX_INDEX] = 1 + //orderCapa.OrderSupport[TS_NEG_DRAWNINEGRID_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_SAVEBITMAP_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_POLYGON_SC_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_POLYGON_CB_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_ELLIPSE_SC_INDEX] = 1 + orderCapa.OrderSupport[TS_NEG_ELLIPSE_CB_INDEX] = 1*/ + //orderCapa.OrderSupport[TS_NEG_FAST_GLYPH_INDEX] = 1 + + inputCapa := c.clientCapabilities[CAPSTYPE_INPUT].(*InputCapability) + inputCapa.Flags = INPUT_FLAG_SCANCODES | INPUT_FLAG_MOUSEX | INPUT_FLAG_UNICODE | + INPUT_FLAG_FASTPATH_INPUT | INPUT_FLAG_FASTPATH_INPUT2 + inputCapa.KeyboardLayout = c.clientCoreData.KbdLayout + inputCapa.KeyboardType = c.clientCoreData.KeyboardType + inputCapa.KeyboardSubType = c.clientCoreData.KeyboardSubType + inputCapa.KeyboardFunctionKey = c.clientCoreData.KeyboardFnKeys + inputCapa.ImeFileName = c.clientCoreData.ImeFileName + + glyphCapa := c.clientCapabilities[CAPSTYPE_GLYPHCACHE].(*GlyphCapability) + /*glyphCapa.GlyphCache[0] = cacheEntry{254, 4} + glyphCapa.GlyphCache[1] = cacheEntry{254, 4} + glyphCapa.GlyphCache[2] = cacheEntry{254, 8} + glyphCapa.GlyphCache[3] = cacheEntry{254, 8} + glyphCapa.GlyphCache[4] = cacheEntry{254, 16} + glyphCapa.GlyphCache[5] = cacheEntry{254, 32} + glyphCapa.GlyphCache[6] = cacheEntry{254, 64} + glyphCapa.GlyphCache[7] = cacheEntry{254, 128} + glyphCapa.GlyphCache[8] = cacheEntry{254, 256} + glyphCapa.GlyphCache[9] = cacheEntry{64, 2048} + glyphCapa.FragCache = 0x01000100*/ + glyphCapa.SupportLevel = GLYPH_SUPPORT_NONE + + // 位图缓存 v2 单元(MS-RDPBCGR 2.2.7.1.8):单元尺寸取 mstsc 量级。 + // CacheFlags 持久键支持位留待 6.4b M2(配合 PERSISTENT_KEY_LIST)。 + rev2, ok := c.clientCapabilities[CAPSTYPE_BITMAPCACHE_REV2].(*BitmapCacheRev2Capability) + if !ok || rev2 == nil { + rev2 = &BitmapCacheRev2Capability{} + c.clientCapabilities[CAPSTYPE_BITMAPCACHE_REV2] = rev2 + } + rev2.CacheCells = [5]BitmapCacheV2CellInfo{ + {NumEntries: 600, MaxCellSize: 256}, + {NumEntries: 1024, MaxCellSize: 512}, + {NumEntries: 4096, MaxCellSize: 1024}, + {}, {}, + } + // v1 位图缓存 caps(0x0004)与 Rev2 一并广告——与 mstsc 行为对齐, + // 部分服务器以 v1 caps 存在与否决定是否启用缓存订单。 + v1, ok := c.clientCapabilities[CAPSTYPE_BITMAPCACHE].(*BitmapCacheCapability) + if !ok || v1 == nil { + v1 = &BitmapCacheCapability{} + c.clientCapabilities[CAPSTYPE_BITMAPCACHE] = v1 + } + + pdu.SharedId = c.sharedId + for _, v := range c.clientCapabilities { + slog.Debug("clientCaps", "type", v.Type(), "value", v) + pdu.CapabilitySets = append(pdu.CapabilitySets, v) + } + pdu.NumberCapabilities = uint16(len(pdu.CapabilitySets)) + pdu.LengthSourceDescriptor = c.demandActivePDU.LengthSourceDescriptor + pdu.SourceDescriptor = c.demandActivePDU.SourceDescriptor + pdu.LengthCombinedCapabilities = c.demandActivePDU.LengthCombinedCapabilities + + c.sendPDU(pdu) +} + +func (c *Client) sendClientFinalizeSynchronizePDU() { + c.sendDataPDU(NewSynchronizeDataPDU(c.channelId)) + c.sendDataPDU(&ControlDataPDU{Action: CTRLACTION_COOPERATE}) + c.sendDataPDU(&ControlDataPDU{Action: CTRLACTION_REQUEST_CONTROL}) + c.maybeSendPersistentKeyList() + c.sendDataPDU(&FontListDataPDU{ListFlags: 0x0003, EntrySize: 0x0032}) +} + +func (c *Client) recvServerSynchronizePDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + pdu, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvServerSynchronizePDU", "err", err) + return + } + dataPdu, ok := pdu.Message.(*DataPDU) + if !ok || dataPdu.Header.PDUType2 != PDUTYPE2_SYNCHRONIZE { + if ok { + slog.Error("recvServerSynchronizePDU ignore datapdu", "type2", dataPdu.Header.PDUType2) + } else { + slog.Error("recvServerSynchronizePDU ignore message", "type", pdu.ShareCtrlHeader.PDUType) + } + slog.Debug("recvServerSynchronizePDU dataPdu", "pdu", &dataPdu) + c.transport.Once("data", c.recvServerSynchronizePDU) + return + } + c.transport.Once("data", c.recvServerControlCooperatePDU) +} + +func (c *Client) recvServerControlCooperatePDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + pdu, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvServerControlCooperatePDU", "err", err) + return + } + dataPdu, ok := pdu.Message.(*DataPDU) + if !ok || dataPdu.Header.PDUType2 != PDUTYPE2_CONTROL { + if ok { + slog.Error("recvServerControlCooperatePDU ignore datapdu", "type2", dataPdu.Header.PDUType2) + } else { + slog.Error("recvServerControlCooperatePDU ignore message", "type", pdu.ShareCtrlHeader.PDUType) + } + c.transport.Once("data", c.recvServerControlCooperatePDU) + return + } + if dataPdu.Data.(*ControlDataPDU).Action != CTRLACTION_COOPERATE { + slog.Error("recvServerControlCooperatePDU ignore", "action", dataPdu.Data.(*ControlDataPDU).Action) + c.transport.Once("data", c.recvServerControlCooperatePDU) + return + } + c.transport.Once("data", c.recvServerControlGrantedPDU) +} + +func (c *Client) recvServerControlGrantedPDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + pdu, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvServerControlGrantedPDU", "err", err) + return + } + dataPdu, ok := pdu.Message.(*DataPDU) + if !ok || dataPdu.Header.PDUType2 != PDUTYPE2_CONTROL { + if ok { + slog.Error("recvServerControlGrantedPDU ignore datapdu", "type2", dataPdu.Header.PDUType2) + } else { + slog.Error("recvServerControlGrantedPDU ignore message", "type", pdu.ShareCtrlHeader.PDUType) + } + c.transport.Once("data", c.recvServerControlGrantedPDU) + return + } + if dataPdu.Data.(*ControlDataPDU).Action != CTRLACTION_GRANTED_CONTROL { + slog.Error("recvServerControlGrantedPDU ignore", "action", dataPdu.Data.(*ControlDataPDU).Action) + c.transport.Once("data", c.recvServerControlGrantedPDU) + return + } + c.transport.Once("data", c.recvServerFontMapPDU) +} + +func (c *Client) recvServerFontMapPDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + pdu, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvServerFontMapPDU", "err", err) + return + } + dataPdu, ok := pdu.Message.(*DataPDU) + if !ok || dataPdu.Header.PDUType2 != PDUTYPE2_FONTMAP { + if ok { + slog.Error("recvServerFontMapPDU ignore datapdu", "type2", dataPdu.Header.PDUType2) + } else { + slog.Error("recvServerFontMapPDU ignore message", "type", pdu.ShareCtrlHeader.PDUType) + } + return + } + c.transport.On("data", c.recvPDU) + + // Tell the server we're ready to receive display updates (MS-RDPBCGR 2.2.11.3.1) + slog.Debug("Sending SuppressOutput (ALLOW_DISPLAY_UPDATES)") + c.sendDataPDU(&SuppressOutputPDU{ + AllowDisplayUpdates: 1, + Right: c.clientCoreData.DesktopWidth - 1, + Bottom: c.clientCoreData.DesktopHeight - 1, + }) + + c.Emit("ready") +} + +func (c *Client) recvPDU(s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + if r.Len() > 0 { + p, err := readPDU(r, c.mppc) + if err != nil { + slog.Error("recvPDU", "err", err) + return + } + if p.ShareCtrlHeader.PDUType == PDUTYPE_DEACTIVATEALLPDU { + // Server is reactivating the session (e.g. desktop resize). + // Signal callers to pause input until "ready" fires again. + slog.Debug("received DeactivateAllPDU during active session, waiting for reactivation") + c.Emit("deactivateAll") + c.transport.Once("data", c.recvDemandActivePDU) + } else if p.ShareCtrlHeader.PDUType == PDUTYPE_SERVER_REDIR_PKT { + if redir, ok := p.Message.(*ServerRedirectionPDU); ok { + c.Emit("redirect", redir) + } + } else if p.ShareCtrlHeader.PDUType == PDUTYPE_DATAPDU { + d := p.Message.(*DataPDU) + if d.Header.PDUType2 == PDUTYPE2_UPDATE { + up := d.Data.(*UpdateDataPDU) + p := up.Udata + if up.UpdateType == FASTPATH_UPDATETYPE_BITMAP { + c.Emit("bitmap", p.(*BitmapUpdateDataPDU).Rectangles) + } else if up.UpdateType == FASTPATH_UPDATETYPE_ORDERS { + c.Emit("orders", p.(*FastPathOrdersPDU).OrderPdus) + } + } else if d.Header.PDUType2 == PDUTYPE2_POINTER { + pp := d.Data.(*PointerDataPDU) + if pp.Pdata != nil { + switch pp.MessageType { + case TS_PTRUPDATE_TYPE_CACHED: + c.Emit("pointer_cached", pp.Pdata.(*FastPathUpdateCachedPDU).CacheIdx) + case TS_PTRUPDATE_TYPE_POINTER: + c.Emit("pointer_update", pp.Pdata.(*FastPathUpdatePointerPDU)) + } + } + if pp.MessageType == TS_PTRUPDATE_TYPE_SYSTEM { + c.Emit("pointer_hide") + } + } else if d.Header.PDUType2 == PDUTYPE2_SET_ERROR_INFO_PDU { + // 服务器在断开连接前会发送错误码,不解析就无法定位断连原因 + ei := d.Data.(*ErrorInfoDataPDU) + LastServerErrorInfo.Store(ei.ErrorInfo) + slog.Warn("server SET_ERROR_INFO", "code", fmt.Sprintf("0x%08X", ei.ErrorInfo)) + } + } + } +} + +func (c *Client) RecvFastPath(secFlag byte, s []byte) { + r := readerPool.Get().(*bytes.Reader) + r.Reset(s) + defer readerPool.Put(r) + for r.Len() > 0 { + updateHeader, err := core.ReadUInt8(r) + if err != nil { + return + } + updateCode := updateHeader & 0x0f + fragmentation := updateHeader & 0x30 + // compression 占 header 的高 2 位:字段值 0b10(左移后 0x80)表示 + // 后随 1 字节 compressionFlags。此前误拿位掩码值与字段值 0x2 直接 + // 比较,恒为假 → 压缩标志字节从未被消费,整条 fast-path 流错位。 + compression := updateHeader & 0xC0 + + var compressionFlags uint8 = 0 + if compression == FASTPATH_OUTPUT_COMPRESSION_USED<<6 { + compressionFlags, err = core.ReadUInt8(r) + if err != nil { + return + } + } + + size, err := core.ReadUint16LE(r) + if err != nil { + return + } + + // Read exactly `size` bytes for this update's payload. + payload, err := core.ReadBytes(int(size), r) + if err != nil { + return + } + + slog.Debug("RecvFastPath", "Code", FastPathUpdateType(updateCode), + "compressionFlags", compressionFlags, + "fragmentation", fragmentation, + "size", size) + + // 先解压再重组:压缩以每个 fast-path 分片为独立单位(各分片头部 + // 携带各自的 compressionFlags),共享同一条 MPPC 历史流;拼接压缩 + // 字节后一次性解压必然产生垃圾。 + decompressed := payload + decFailed := false + if compressionFlags != 0 && c.mppc != nil { + out, err := c.mppc.Decompress(compressionFlags, payload) + if err != nil { + decFailed = true + if c.decErrLogs < 4 { + c.decErrLogs++ + slog.Warn("RecvFastPath: MPPC decompression failed", "err", err, + "cf", fmt.Sprintf("%02X", compressionFlags), + "frag", fragmentation, "size", size) + } + } else { + decompressed = out + } + } + + // 分片重组:续片的 updateCode 不可信(MS-RDPBCGR 2.2.9.1.1.3.1 + // 规定忽略),解析时使用首片记住的 code。解压失败的分片会把所在 + // 更新标记为中毒:吞掉其剩余分片,避免孤儿分片拼出垃圾更新。 + if fragmentation != FASTPATH_FRAGMENT_SINGLE { + if fragmentation == FASTPATH_FRAGMENT_FIRST { + c.buff.Reset() + c.fragCode = updateCode + c.fragActive = true + c.fragPoisoned = decFailed + } else if !c.fragActive || c.fragPoisoned { + continue // 无首片的孤儿分片 / 中毒更新的剩余分片 + } + c.buff.Write(decompressed) + if fragmentation != FASTPATH_FRAGMENT_LAST { + continue + } + c.fragActive = false + if c.fragPoisoned { + continue + } + payload = c.buff.Bytes() + updateCode = c.fragCode + } else if decFailed { + continue + } else { + // 单片更新:解析对象必须是解压后的结果(此前漏赋值导致 + // 嗅探/解析拿到原始压缩数据,单片位图全部报废)。 + payload = decompressed + } + + // Surface Commands:Win10 会把一条表面命令拆到多个 fast-path PDU + //(实测 frag 位 FIRST/NEXT/LAST 使用规范),按命令流自定界累积解析,命令完整即发射。 + if updateCode == FASTPATH_UPDATETYPE_SURFCMDS { + if decFailed { + continue + } + c.surfBuf = append(c.surfBuf, payload...) + for { + result, consumed, needMore, valid, dropped, hdrLen := parseSurfaceCommandsIncremental(c.surfBuf, c.surfHdrLen) + if c.surfHdrLen == 0 && hdrLen != 0 { + c.surfHdrLen = hdrLen // 首条命令锁定头部布局(会话内恒定) + } + if !valid { + if c.surfResetDumped < 2 { + c.surfResetDumped++ + slog.Warn("SURFCMDS stream reset full dump", + "n", c.surfResetDumped, "bufLen", len(c.surfBuf), + "hdrLen", c.surfHdrLen, + "hex", fmt.Sprintf("% X", c.surfBuf)) + } + slog.Warn("SURFCMDS stream reset", + "bufLen", len(c.surfBuf), + "hex", fmt.Sprintf("% X", c.surfBuf[:min(24, len(c.surfBuf))])) + c.surfBuf = c.surfBuf[:0] + c.surfHdrLen = 0 // 布局判定作废,下批重新锁定 + break + } + if len(result.Rects) > 0 { + c.Emit("bitmap", result.Rects) + } + for _, fid := range result.FrameIDs { + c.sendDataPDU(&FrameAcknowledgeDataPDU{FrameID: fid}) + } + c.surfBuf = c.surfBuf[consumed:] + if dropped { + // 尾部未知结构:已解析前缀照常上屏,仅丢弃尾巴。 + // 实测 Win10 整屏重绘批次末尾带 11 字节未文档化尾巴, + // 整批报废会造成周期性花屏+断流。 + slog.Debug("SURFCMDS tail dropped", "consumed", consumed, "bufLen", len(c.surfBuf)) + c.surfBuf = c.surfBuf[:0] + break + } + if needMore || consumed == 0 || len(c.surfBuf) < 2 { + break + } + } + if len(c.surfBuf) > 8<<20 { + c.surfBuf = c.surfBuf[:0] // 防御:异常累积封顶 + } + continue + } + + // fast-path 位图:Win10 会把一个位图更新拆到多个 fast-path PDU + //(首段携带 updateType + numberRectangles + 前几个矩形,后续段 + // 继续补充矩形数据)。将解压结果持续累积,一旦能完整解析出全部 + // 矩形就发射。非位图更新仍走通用分片状态机。 + // 同样必须追加 payload(重组缓冲)而非 decompressed,理由同上。 + if updateCode == FASTPATH_UPDATETYPE_BITMAP { + if decFailed { + continue + } + c.bmpFrame = append(c.bmpFrame, payload...) + for { + rects, consumed, complete, valid := parseBitmapFrame(c.bmpFrame) + if !valid { + slog.Warn("BITMAP frame reset", + "bufLen", len(c.bmpFrame), + "hex", fmt.Sprintf("% X", c.bmpFrame[:min(24, len(c.bmpFrame))])) + c.bmpFrame = c.bmpFrame[:0] + break + } + if !complete { + break + } + if len(rects) > 0 { + if c.bmpRectLogs < 24 { + for i, rc := range rects { + if c.bmpRectLogs >= 24 { + break + } + c.bmpRectLogs++ + slog.Debug("BITMAP rect", "i", i, "n", len(rects), + "dx", rc.DestLeft, "dy", rc.DestTop, + "dr", rc.DestRight, "db", rc.DestBottom, + "w", rc.Width, "h", rc.Height, + "bpp", rc.BitsPerPixel, "flags", fmt.Sprintf("0x%04X", rc.Flags), + "len", len(rc.BitmapDataStream)) + } + } + c.Emit("bitmap", rects) + } + c.bmpFrame = c.bmpFrame[consumed:] + if len(c.bmpFrame) < 4 { + break + } + } + continue + } + + if updateCode == FASTPATH_UPDATETYPE_POINTER { + // 取证:指针更新原始字节,核对 xorBpp 后各字段对齐。 + slog.Debug("PTR raw", "hex", fmt.Sprintf("% X", payload[:min(32, len(payload))])) + } + + pr := bytes.NewReader(payload) + p, err := readFastPathUpdatePDU(pr, updateCode) + if err != nil { + slog.Warn("readFastPathUpdatePDU:", "Code", FastPathUpdateType(updateCode), "err", err) + continue + } + + if updateCode == FASTPATH_UPDATETYPE_BITMAP { + c.Emit("bitmap", p.Data.(*FastPathBitmapUpdateDataPDU).Rectangles) + } else if updateCode == FASTPATH_UPDATETYPE_COLOR { + c.Emit("color", p.Data.(*FastPathColorPdu)) + } else if updateCode == FASTPATH_UPDATETYPE_ORDERS { + c.Emit("orders", p.Data.(*FastPathOrdersPDU).OrderPdus) + } else if updateCode == FASTPATH_UPDATETYPE_PTR_NULL { + c.Emit("pointer_hide") + } else if updateCode == FASTPATH_UPDATETYPE_PTR_DEFAULT { + c.Emit("pointer_default") + } else if updateCode == FASTPATH_UPDATETYPE_PTR_POSITION { + pp := p.Data.(*FastPathPointerPositionPDU) + c.Emit("pointer_position", pp.X, pp.Y) + } else if updateCode == FASTPATH_UPDATETYPE_CACHED { + c.Emit("pointer_cached", p.Data.(*FastPathUpdateCachedPDU).CacheIdx) + } else if updateCode == FASTPATH_UPDATETYPE_POINTER { + c.Emit("pointer_update", p.Data.(*FastPathUpdatePointerPDU)) + } + } +} + +type InputEventsInterface interface { + Serialize() []byte +} + +// fastPathEncoder is implemented by input event types that know how to +// produce their Fast-Path Input wire encoding (MS-RDPBCGR §2.2.8.1.2.2). +type fastPathEncoder interface { + FastPathEncode(buf []byte) []byte +} + +func (c *Client) SendInputEvents(msgType uint16, events []InputEventsInterface) { + if c.serverFastPathInput && c.fastPathSender != nil && c.canSendFastPathInput(events) { + if c.sendFastPathInputEvents(events) { + return + } + // Fall back to slow-path on send failure (e.g. legacy encryption). + } + + p := &ClientInputEventPDU{} + p.NumEvents = uint16(len(events)) + p.SlowPathInputEvents = make([]SlowPathInputEvent, 0, p.NumEvents) + for _, in := range events { + seria := in.Serialize() + s := SlowPathInputEvent{0, msgType, len(seria), seria} + p.SlowPathInputEvents = append(p.SlowPathInputEvents, s) + } + + c.sendDataPDU(p) +} + +// canSendFastPathInput reports whether every event in the batch implements +// the fast-path encoder. Falls back to slow-path if any event type doesn't +// (currently just SynchronizeEvent, which the client never sends). +func (c *Client) canSendFastPathInput(events []InputEventsInterface) bool { + if len(events) == 0 || len(events) > 15 { + return false + } + for _, e := range events { + if _, ok := e.(fastPathEncoder); !ok { + return false + } + } + return true +} + +func (c *Client) sendFastPathInputEvents(events []InputEventsInterface) bool { + buf := fastPathBufPool.Get().([]byte) + buf = buf[:0] + buf = append(buf, byte(len(events))) + for _, e := range events { + buf = e.(fastPathEncoder).FastPathEncode(buf) + } + _, err := c.fastPathSender.SendFastPath(0, buf) + fastPathBufPool.Put(buf[:cap(buf)]) + if err != nil { + // Disable for the rest of the session so we don't keep paying the + // failed-attempt cost on every input event. + c.serverFastPathInput = false + slog.Warn("fast-path input disabled, falling back to slow-path", "err", err) + return false + } + return true +} + +// SendRefreshRect requests the server to redraw the given screen rectangle. +// This causes the server to send a full refresh (including a new H.264 IDR) +// for the specified region, which is useful after a decoder reset. +func (c *Client) SendRefreshRect(width, height uint16) { + slog.Debug("PDU: SendRefreshRect", "w", width, "h", height) + c.sendDataPDU(&RefreshRectPDU{ + NumberOfAreas: 1, + Right: width - 1, + Bottom: height - 1, + }) +} + +// SendForceRefresh asks the server for a complete display repaint by toggling +// SuppressOutput off→on. Per MS-RDPBCGR 2.2.11.3.1, sending ALLOW_DISPLAY_UPDATES +// after SUPPRESS_DISPLAY_UPDATES forces the server to send a fresh full-screen +// update — for the RDPGFX H.264 pipeline this means a new IDR frame, which is +// what we need to recover after a hardware-decoder hard reset. Plain +// SendRefreshRect is sometimes silently ignored by Windows servers while a +// video stream is active; this is the reliable fallback used by mstsc/FreeRDP. +func (c *Client) SendForceRefresh(width, height uint16) { + slog.Debug("PDU: SendForceRefresh (suppress→allow)", "w", width, "h", height) + c.sendDataPDU(&SuppressOutputPDU{ + AllowDisplayUpdates: 0x00, // SUPPRESS_DISPLAY_UPDATES + }) + c.sendDataPDU(&SuppressOutputPDU{ + AllowDisplayUpdates: 0x01, // ALLOW_DISPLAY_UPDATES + Right: width - 1, + Bottom: height - 1, + }) +} + +// bitmapDropLogs:垃圾帧丢弃诊断日志限次计数器。 +var bitmapDropLogs = 0 + +// parseBitmapFrame 尝试把累积缓冲解析为完整的 fast-path 位图更新 +// (updateType(2) + numberRectangles(2) + 矩形数组)。 +// 返回 valid=false 表示缓冲开头不是合法位图结构(应丢弃); +// complete=false 表示数据不足(应继续累积); +// complete=true 时 consumed 为本次更新占用的字节数,尾部可能残留 +// 下一更新的开头。 +func parseBitmapFrame(buf []byte) (rects []BitmapData, consumed int, complete bool, valid bool) { + if len(buf) < 4 { + return nil, 0, false, true + } + if u16 := binary.LittleEndian.Uint16(buf[0:2]); u16 != 0x0001 { // UPDATE_TYPE_BITMAP + return nil, 0, false, false + } + nr := int(binary.LittleEndian.Uint16(buf[2:4])) + if nr > 4096 { + if bitmapDropLogs < 3 { + bitmapDropLogs++ + slog.Warn("bitmap frame dropped: implausible rect count", + "nr", nr, "len", len(buf), + "prefix", fmt.Sprintf("% X", buf[:min(24, len(buf))])) + } + return nil, 0, false, false + } + rects = make([]BitmapData, 0, nr) + pos := 4 + for i := 0; i < nr; i++ { + if pos+18 > len(buf) { + return nil, 0, false, true + } + r := BitmapData{} + r.DestLeft = binary.LittleEndian.Uint16(buf[pos:]) + r.DestTop = binary.LittleEndian.Uint16(buf[pos+2:]) + r.DestRight = binary.LittleEndian.Uint16(buf[pos+4:]) + r.DestBottom = binary.LittleEndian.Uint16(buf[pos+6:]) + r.Width = binary.LittleEndian.Uint16(buf[pos+8:]) + r.Height = binary.LittleEndian.Uint16(buf[pos+10:]) + r.BitsPerPixel = binary.LittleEndian.Uint16(buf[pos+12:]) + r.Flags = binary.LittleEndian.Uint16(buf[pos+14:]) + bl := int(binary.LittleEndian.Uint16(buf[pos+16:])) + hdr := 0 + if r.Flags&BITMAP_COMPRESSION != 0 && r.Flags&NO_BITMAP_COMPRESSION_HDR == 0 { + if pos+18+8 > len(buf) { + return nil, 0, false, true + } + bl = int(binary.LittleEndian.Uint16(buf[pos+18+2 : pos+18+4])) + hdr = 8 + } + if bl < 0 || pos+18+hdr+bl > len(buf) { + return nil, 0, false, true + } + r.BitmapDataStream = buf[pos+18+hdr : pos+18+hdr+bl] + pos += 18 + hdr + bl + rects = append(rects, r) + } + return rects, pos, true, true +} + +// SetPersistentKeyList 注册客户端持久位图缓存持有的键(bitmap 管线)。 +// 须在 DemandActive 之前调用;finalize 时仅当服务器广告了 +// CAPSTYPE_BITMAPCACHE_HOSTSUPPORT 且键非空才实际发送。 +func (p *PDULayer) SetPersistentKeyList(keys []uint64) { + p.bitmapCachePersistKeys = keys +} + +// maybeSendPersistentKeyList 在 Connection Finalization 序列 +// (REQUEST_CONTROL 之后、FONT_LIST 之前)发送持久缓存键列表, +// 位置与 FreeRDP/MS-RDPBCGR §2.2.1.17 一致。 +func (p *PDULayer) maybeSendPersistentKeyList() { + if len(p.bitmapCachePersistKeys) == 0 { + return + } + hs, ok := p.serverCapabilities[CAPSTYPE_BITMAPCACHE_HOSTSUPPORT].(*BitmapCacheHostSupportCapability) + if !ok || hs.CacheVersion < 1 { + slog.Info("bmpcache: server lacks bitmap cache host support, skip persistent key list") + return + } + var cells [5]uint16 + if rev2, ok := p.clientCapabilities[CAPSTYPE_BITMAPCACHE_REV2].(*BitmapCacheRev2Capability); ok { + for i := range rev2.CacheCells { + if i < len(cells) { + cells[i] = rev2.CacheCells[i].NumEntries + } + } + } + keys := p.bitmapCachePersistKeys + if len(keys) > 2042 { + // FreeRDP 同款上限:>2042 条会触发服务器错误 + keys = keys[:2042] + slog.Info("bmpcache: truncating persistent key list", "kept", len(keys)) + } + pdu := NewPDU(p.userId, NewPersistentKeyListPDU(p.sharedId, keys, cells)) + p.transport.Write(pdu.serialize()) + slog.Info("bmpcache: sent persistent key list", "keys", len(keys)) +} diff --git a/protocol/pdu/ycbcr_amd64.go b/protocol/pdu/ycbcr_amd64.go new file mode 100644 index 0000000..1e0419d --- /dev/null +++ b/protocol/pdu/ycbcr_amd64.go @@ -0,0 +1,34 @@ +//go:build amd64 + +package pdu + +import "unsafe" + +// ycoCgToBGRANoSub converts Y, Co, Cg planes to BGRA using SSE2. +// count = total pixel count. shift = colorLossLevel - 1. +// Alpha is always 0xFF (caller handles non-0xFF alpha separately). +func ycoCgToBGRANoSub(pixels []byte, yPlane, coPlane, cgPlane []byte, count int, shift uint8) { + count8 := count &^ 7 + if count8 > 0 { + ycoCgToBGRASSE2( + unsafe.Pointer(&pixels[0]), + unsafe.Pointer(&yPlane[0]), + unsafe.Pointer(&coPlane[0]), + unsafe.Pointer(&cgPlane[0]), + count8, int(shift), + ) + } + for i := count8; i < count; i++ { + yVal := int16(yPlane[i]) + coVal := int16(int8(byte(int16(coPlane[i]) << shift))) + cgVal := int16(int8(byte(int16(cgPlane[i]) << shift))) + off := i * 4 + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + pixels[off+3] = 0xFF + } +} + +//go:noescape +func ycoCgToBGRASSE2(pixels, yPlane, coPlane, cgPlane unsafe.Pointer, count, shift int) diff --git a/protocol/pdu/ycbcr_amd64.s b/protocol/pdu/ycbcr_amd64.s new file mode 100644 index 0000000..07a0104 --- /dev/null +++ b/protocol/pdu/ycbcr_amd64.s @@ -0,0 +1,100 @@ +// SSE2 implementation of ycoCgToBGRASSE2. +// Processes 8 pixels per iteration (no chroma subsampling, alpha = 0xFF). +// Stack ABI (ABI0): +// pixels+0(FP) unsafe.Pointer (dst, BGRA output) +// yPlane+8(FP) unsafe.Pointer +// coPlane+16(FP) unsafe.Pointer +// cgPlane+24(FP) unsafe.Pointer +// count+32(FP) int (multiple of 8) +// shift+40(FP) int + +#include "textflag.h" + +// func ycoCgToBGRASSE2(pixels, yPlane, coPlane, cgPlane unsafe.Pointer, count, shift int) +TEXT ·ycoCgToBGRASSE2(SB),NOSPLIT,$0-48 + MOVQ pixels+0(FP), DI + MOVQ yPlane+8(FP), SI + MOVQ coPlane+16(FP), BX + MOVQ cgPlane+24(FP), R8 + MOVQ count+32(FP), CX + MOVQ shift+40(FP), AX + + // Build XMM shift count = shift+8 in low 64 bits (used by PSLLW). + ADDQ $8, AX + MOVQ AX, X12 // X12 = shift+8 (PSLLW count register) + + PXOR X7, X7 // X7 = zero vector (for zero-extension) + + SHRQ $3, CX // CX = count/8 (loop iterations) + +loop_yco: + // Load 8 Y bytes; zero-extend each to uint16. + MOVQ (SI), X1 // X1[63:0] = 8 Y bytes + PUNPCKLBW X7, X1 // X1 = uint16[0..7] Y values (0-255) + + // Load 8 Co bytes; zero-extend then sign-extend with shift: + // coVal = int16(int8(byte(uint16(co) << shift))) + // = (uint16(co) << (shift+8)) >> 8 [arithmetic] + MOVQ (BX), X2 + PUNPCKLBW X7, X2 // X2 = uint16 co values + PSLLW X12, X2 // X2 <<= shift+8 + PSRAW $8, X2 // X2 = signed int16 coVal + + // Same for Cg. + MOVQ (R8), X3 + PUNPCKLBW X7, X3 + PSLLW X12, X3 + PSRAW $8, X3 // X3 = signed int16 cgVal + + // B = Y - Co - Cg. + MOVO X1, X4 + PSUBW X2, X4 + PSUBW X3, X4 + + // G = Y + Cg. + MOVO X1, X5 + PADDW X3, X5 + + // R = Y + Co - Cg. + MOVO X1, X6 + PADDW X2, X6 + PSUBW X3, X6 + + // Pack B and G to uint8 with unsigned saturation (clamp to [0,255]): + // PACKUSWB dst,src: dst = [sat_u8(dst[0..7]), sat_u8(src[0..7])] + // After: X4 = [B0..B7, G0..G7] + PACKUSWB X5, X4 + + // Pack R and 0xFF-alpha: + // PCMPEQB X10,X10 → all bits 1 = 0xFF per byte. + PCMPEQB X10, X10 // X10 = 0xFF...FF + PACKUSWB X10, X6 // X6 = [R0..R7, FF..FF] + + // Interleave B and G bytes: [B0,G0,B1,G1,...,B7,G7]. + // X4 = [B0..B7 | G0..G7]; shift copy right 8 bytes → [G0..G7 | 0..0]. + MOVO X4, X8 + PSRLDQ $8, X8 // X8 = [G0..G7, 0..0] + PUNPCKLBW X8, X4 // X4 = [B0,G0,B1,G1,...,B7,G7] + + // Interleave R and alpha bytes: [R0,FF,R1,FF,...,R7,FF]. + MOVO X6, X9 + PSRLDQ $8, X9 // X9 = [FF..FF, 0..0] + PUNPCKLBW X9, X6 // X6 = [R0,FF,R1,FF,...,R7,FF] + + // Interleave BG and RA halfwords to produce BGRA dwords: + // PUNPCKLWL: low 4 words → [BG0,RA0,BG1,RA1,BG2,RA2,BG3,RA3] + // = bytes [B0,G0,R0,FF, B1,G1,R1,FF, B2,G2,R2,FF, B3,G3,R3,FF] + MOVO X4, X11 + PUNPCKLWL X6, X11 // X11 = low 4 BGRA pixels + PUNPCKHWL X6, X4 // X4 = high 4 BGRA pixels + + MOVOU X11, (DI) + MOVOU X4, 16(DI) + + ADDQ $8, SI + ADDQ $8, BX + ADDQ $8, R8 + ADDQ $32, DI + DECQ CX + JNZ loop_yco + RET diff --git a/protocol/pdu/ycbcr_arm64.go b/protocol/pdu/ycbcr_arm64.go new file mode 100644 index 0000000..67bbcbf --- /dev/null +++ b/protocol/pdu/ycbcr_arm64.go @@ -0,0 +1,31 @@ +//go:build arm64 + +package pdu + +import "unsafe" + +func ycoCgToBGRANoSub(pixels []byte, yPlane, coPlane, cgPlane []byte, count int, shift uint8) { + count8 := count &^ 7 + if count8 > 0 { + ycoCgToBGRANEON( + unsafe.Pointer(&pixels[0]), + unsafe.Pointer(&yPlane[0]), + unsafe.Pointer(&coPlane[0]), + unsafe.Pointer(&cgPlane[0]), + count8, int(shift), + ) + } + for i := count8; i < count; i++ { + yVal := int16(yPlane[i]) + coVal := int16(int8(byte(int16(coPlane[i]) << shift))) + cgVal := int16(int8(byte(int16(cgPlane[i]) << shift))) + off := i * 4 + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + pixels[off+3] = 0xFF + } +} + +//go:noescape +func ycoCgToBGRANEON(pixels, yPlane, coPlane, cgPlane unsafe.Pointer, count, shift int) diff --git a/protocol/pdu/ycbcr_arm64.s b/protocol/pdu/ycbcr_arm64.s new file mode 100644 index 0000000..9a53a6c --- /dev/null +++ b/protocol/pdu/ycbcr_arm64.s @@ -0,0 +1,79 @@ +// Scalar ARM64 implementation of ycoCgToBGRANEON. +// Processes one pixel per iteration (no chroma subsampling, alpha = 0xFF). +// Go arm64 assembler lacks SSHR (signed vector shift right), SSHL (by register), +// and SQXTUN (saturating narrow), so we use scalar GP instructions instead. +// Stack ABI (ABI0): +// pixels+0(FP) unsafe.Pointer +// yPlane+8(FP) unsafe.Pointer +// coPlane+16(FP) unsafe.Pointer +// cgPlane+24(FP) unsafe.Pointer +// count+32(FP) int +// shift+40(FP) int + +#include "textflag.h" + +// func ycoCgToBGRANEON(pixels, yPlane, coPlane, cgPlane unsafe.Pointer, count, shift int) +TEXT ·ycoCgToBGRANEON(SB),NOSPLIT,$0-48 + MOVD pixels+0(FP), R0 // dst + MOVD yPlane+8(FP), R1 // Y plane + MOVD coPlane+16(FP), R2 // Co plane + MOVD cgPlane+24(FP), R3 // Cg plane + MOVD count+32(FP), R4 // pixel count + MOVD shift+40(FP), R5 // shift amount + + CBZ R4, done + + MOVD $0, R12 // const 0 + MOVD $255, R13 // const 255 + +loop: + MOVBU (R1), R6 // y = Y[i] + MOVBU (R2), R7 // co = Co[i] + MOVBU (R3), R8 // cg = Cg[i] + ADD $1, R1 + ADD $1, R2 + ADD $1, R3 + + // coVal = int8(co << shift): shift left, then sign-extend low byte. + LSL R5, R7, R7 // R7 = co << shift (low 8 bits hold result) + SXTB R7, R7 // R7 = int64(int8(R7)) + + // cgVal = int8(cg << shift) + LSL R5, R8, R8 + SXTB R8, R8 + + // B = clamp(y - co - cg) + SUB R7, R6, R9 // R9 = y - co + SUB R8, R9, R9 // R9 = y - co - cg + CMP R12, R9 + CSEL LT, R12, R9, R9 + CMP R13, R9 + CSEL GT, R13, R9, R9 + + // G = clamp(y + cg) + ADD R8, R6, R10 // R10 = y + cg + CMP R12, R10 + CSEL LT, R12, R10, R10 + CMP R13, R10 + CSEL GT, R13, R10, R10 + + // R = clamp(y + co - cg) + ADD R7, R6, R11 // R11 = y + co + SUB R8, R11, R11 // R11 = y + co - cg + CMP R12, R11 + CSEL LT, R12, R11, R11 + CMP R13, R11 + CSEL GT, R13, R11, R11 + + // Store BGRA pixel. + MOVBU R9, 0(R0) + MOVBU R10, 1(R0) + MOVBU R11, 2(R0) + MOVBU R13, 3(R0) // alpha = 255 + ADD $4, R0 + + SUBS $1, R4, R4 + BNE loop + +done: + RET diff --git a/protocol/pdu/ycbcr_generic.go b/protocol/pdu/ycbcr_generic.go new file mode 100644 index 0000000..2dc5487 --- /dev/null +++ b/protocol/pdu/ycbcr_generic.go @@ -0,0 +1,55 @@ +//go:build !amd64 && !arm64 + +package pdu + +// ycoCgToBGRANoSub converts count pixels from YCoCg planes to interleaved BGRA, +// assuming no chroma subsampling and no alpha plane override. +// pixels must have capacity >= count*4. +func ycoCgToBGRANoSub(pixels, yPlane, coPlane, cgPlane []byte, count int, shift uint8) { + i := 0 + for ; i+4 <= count; i += 4 { + off := i * 4 + + yVal := int16(yPlane[i]) + coVal := int16(int8(byte(int16(coPlane[i]) << shift))) + cgVal := int16(int8(byte(int16(cgPlane[i]) << shift))) + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + pixels[off+3] = 0xFF + + yVal = int16(yPlane[i+1]) + coVal = int16(int8(byte(int16(coPlane[i+1]) << shift))) + cgVal = int16(int8(byte(int16(cgPlane[i+1]) << shift))) + pixels[off+4] = clampByte(yVal - coVal - cgVal) + pixels[off+5] = clampByte(yVal + cgVal) + pixels[off+6] = clampByte(yVal + coVal - cgVal) + pixels[off+7] = 0xFF + + yVal = int16(yPlane[i+2]) + coVal = int16(int8(byte(int16(coPlane[i+2]) << shift))) + cgVal = int16(int8(byte(int16(cgPlane[i+2]) << shift))) + pixels[off+8] = clampByte(yVal - coVal - cgVal) + pixels[off+9] = clampByte(yVal + cgVal) + pixels[off+10] = clampByte(yVal + coVal - cgVal) + pixels[off+11] = 0xFF + + yVal = int16(yPlane[i+3]) + coVal = int16(int8(byte(int16(coPlane[i+3]) << shift))) + cgVal = int16(int8(byte(int16(cgPlane[i+3]) << shift))) + pixels[off+12] = clampByte(yVal - coVal - cgVal) + pixels[off+13] = clampByte(yVal + cgVal) + pixels[off+14] = clampByte(yVal + coVal - cgVal) + pixels[off+15] = 0xFF + } + for ; i < count; i++ { + yVal := int16(yPlane[i]) + coVal := int16(int8(byte(int16(coPlane[i]) << shift))) + cgVal := int16(int8(byte(int16(cgPlane[i]) << shift))) + off := i * 4 + pixels[off] = clampByte(yVal - coVal - cgVal) + pixels[off+1] = clampByte(yVal + cgVal) + pixels[off+2] = clampByte(yVal + coVal - cgVal) + pixels[off+3] = 0xFF + } +} diff --git a/protocol/sec/sec.go b/protocol/sec/sec.go new file mode 100644 index 0000000..1d0f8d5 --- /dev/null +++ b/protocol/sec/sec.go @@ -0,0 +1,994 @@ +package sec + +import ( + "bytes" + "crypto/md5" + "crypto/rand" + "crypto/rc4" + "crypto/rsa" + "crypto/sha1" + "encoding/binary" + "errors" + "fmt" + "github.com/lunixbochs/struc" + "io" + "log/slog" + "unicode/utf16" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" + "git.zeroonesoft.cn/golib/rdplib/protocol/lic" + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/gcc" +) + +// Pre-computed padding bytes used in MAC generation (avoids per-call allocations). +var ( + macPad36 [40]byte + macPad5C [48]byte +) + +func init() { + for i := range macPad36 { + macPad36[i] = 0x36 + } + for i := range macPad5C { + macPad5C[i] = 0x5c + } +} + +/** + * SecurityFlag + * @see http://msdn.microsoft.com/en-us/library/cc240579.aspx + */ +const ( + EXCHANGE_PKT uint16 = 0x0001 + TRANSPORT_REQ = 0x0002 + TRANSPORT_RSP = 0x0004 + ENCRYPT = 0x0008 + RESET_SEQNO = 0x0010 + IGNORE_SEQNO = 0x0020 + INFO_PKT = 0x0040 + LICENSE_PKT = 0x0080 + LICENSE_ENCRYPT_CS = 0x0200 + LICENSE_ENCRYPT_SC = 0x0200 + REDIRECTION_PKT = 0x0400 + SECURE_CHECKSUM = 0x0800 + AUTODETECT_REQ = 0x1000 + AUTODETECT_RSP = 0x2000 + HEARTBEAT = 0x4000 + FLAGSHI_VALID = 0x8000 +) + +const ( + INFO_MOUSE uint32 = 0x00000001 + INFO_DISABLECTRLALTDEL = 0x00000002 + INFO_AUTOLOGON = 0x00000008 + INFO_UNICODE = 0x00000010 + INFO_MAXIMIZESHELL = 0x00000020 + INFO_LOGONNOTIFY = 0x00000040 + INFO_COMPRESSION = 0x00000080 + INFO_ENABLEWINDOWSKEY = 0x00000100 + INFO_REMOTECONSOLEAUDIO = 0x00002000 + INFO_FORCE_ENCRYPTED_CS_PDU = 0x00004000 + INFO_RAIL = 0x00008000 + INFO_LOGONERRORS = 0x00010000 + INFO_MOUSE_HAS_WHEEL = 0x00020000 + INFO_PASSWORD_IS_SC_PIN = 0x00040000 + INFO_NOAUDIOPLAYBACK = 0x00080000 + INFO_USING_SAVED_CREDS = 0x00100000 + INFO_AUDIOCAPTURE = 0x00200000 + INFO_VIDEO_DISABLE = 0x00400000 + INFO_CompressionTypeMask = 0x00001E00 + // INFO_CompressionTypeMask 取值(MS-RDPBCGR 2.2.1.11.1.1, + // PACKET_COMPR_TYPE_*:4 位枚举,表示客户端支持的最高压缩包): + // 0x0=RDP4(8K MPPC) 0x1=RDP5(64K MPPC) 0x2=RDP6.0 0x3=RDP6.1 + // (后两者是 MS-RDPEGDI 的独立压缩包,core 包未实现,不可声明) + INFO_CompressionTypeRDP5 = 0x00000200 +) + +const ( + AF_INET uint16 = 0x00002 + AF_INET6 = 0x0017 +) + +const ( + PERF_DISABLE_WALLPAPER uint32 = 0x00000001 + PERF_DISABLE_FULLWINDOWDRAG = 0x00000002 + PERF_DISABLE_MENUANIMATIONS = 0x00000004 + PERF_DISABLE_THEMING = 0x00000008 + PERF_DISABLE_CURSOR_SHADOW = 0x00000020 + PERF_DISABLE_CURSORSETTINGS = 0x00000040 + PERF_ENABLE_FONT_SMOOTHING = 0x00000080 + PERF_ENABLE_DESKTOP_COMPOSITION = 0x00000100 +) + +const ( + FASTPATH_OUTPUT_SECURE_CHECKSUM = 0x1 + FASTPATH_OUTPUT_ENCRYPTED = 0x2 +) + +type ClientAutoReconnect struct { + CbAutoReconnectLen uint16 + CbLen uint32 + Version uint32 + LogonId uint32 + SecVerifier []byte +} + +func NewClientAutoReconnect(id uint32, random []byte) *ClientAutoReconnect { + return &ClientAutoReconnect{ + CbAutoReconnectLen: 28, + CbLen: 28, + Version: 1, + LogonId: id, + SecVerifier: nla.HMAC_MD5(random, random), + } +} + +type RDPExtendedInfo struct { + ClientAddressFamily uint16 `struc:"little"` + CbClientAddress uint16 `struc:"little,sizeof=ClientAddress"` + ClientAddress []byte `struc:"[]byte"` + CbClientDir uint16 `struc:"little,sizeof=ClientDir"` + ClientDir []byte `struc:"[]byte"` + ClientTimeZone []byte `struc:"[172]byte"` + ClientSessionId uint32 `struc:"litttle"` + PerformanceFlags uint32 `struc:"little"` + AutoReconnect *ClientAutoReconnect + // [MS-RDPBCGR] 2.2.1.11.1.1.1 中 cbAutoReconnectCookie 之后的 optional 字段链。 + // Windows 8/2012+ 服务器按 mstsc 的完整布局解析,缺失时报 + // ERRINFO_TIMEZONE_KEY_NAME_LENGTH_TOO_SHORT (0x112F) 并断开连接。 + DynDSTTimeZoneKeyName []byte // UTF-16LE,不含结尾空字符 + DynamicDaylightTimeDisabled uint16 +} + +func NewExtendedInfo(auto *ClientAutoReconnect) *RDPExtendedInfo { + return &RDPExtendedInfo{ + ClientAddressFamily: AF_INET, + ClientAddress: []byte{0, 0}, + ClientDir: []byte{0, 0}, + ClientTimeZone: make([]byte, 172), + ClientSessionId: 0, + // 视觉全开:启用壁纸/整窗拖动/菜单动画(对齐 mstsc 默认体验档), + // 仅保留字体平滑与桌面合成增强 + PerformanceFlags: (PERF_ENABLE_FONT_SMOOTHING | + PERF_ENABLE_DESKTOP_COMPOSITION), + AutoReconnect: auto, + } +} + +func (o *RDPExtendedInfo) Serialize() []byte { + buff := &bytes.Buffer{} + core.WriteUInt16LE(o.ClientAddressFamily, buff) + core.WriteUInt16LE(uint16(len(o.ClientAddress)), buff) + core.WriteBytes(o.ClientAddress, buff) + core.WriteUInt16LE(uint16(len(o.ClientDir)), buff) + core.WriteBytes(o.ClientDir, buff) + core.WriteBytes(o.ClientTimeZone, buff) + core.WriteUInt32LE(o.ClientSessionId, buff) + core.WriteUInt32LE(o.PerformanceFlags, buff) + + // cbAutoReconnectCookie 恒存在(无 cookie 时写 0),与 mstsc/FreeRDP 一致 + if o.AutoReconnect != nil { + core.WriteUInt16LE(o.AutoReconnect.CbAutoReconnectLen, buff) + core.WriteUInt32LE(o.AutoReconnect.CbLen, buff) + core.WriteUInt32LE(o.AutoReconnect.Version, buff) + core.WriteUInt32LE(o.AutoReconnect.LogonId, buff) + core.WriteBytes(o.AutoReconnect.SecVerifier, buff) + } else { + core.WriteUInt16LE(0, buff) + } + + // reserved1、reserved2、dynamic DST 字段链(FreeRDP info.c 同布局) + core.WriteUInt16LE(0, buff) + core.WriteUInt16LE(0, buff) + core.WriteUInt16LE(uint16(len(o.DynDSTTimeZoneKeyName)), buff) + core.WriteBytes(o.DynDSTTimeZoneKeyName, buff) + core.WriteUInt16LE(o.DynamicDaylightTimeDisabled, buff) + + return buff.Bytes() +} + +type RDPInfo struct { + CodePage uint32 + Flag uint32 + CbDomain uint16 + CbUserName uint16 + CbPassword uint16 + CbAlternateShell uint16 + CbWorkingDir uint16 + Domain []byte + UserName []byte + Password []byte + AlternateShell []byte + WorkingDir []byte + ExtendedInfo *RDPExtendedInfo +} + +func NewRDPInfo() *RDPInfo { + info := &RDPInfo{ + Flag: INFO_MOUSE | INFO_UNICODE | INFO_MAXIMIZESHELL | + INFO_ENABLEWINDOWSKEY | INFO_DISABLECTRLALTDEL | INFO_MOUSE_HAS_WHEEL | + INFO_FORCE_ENCRYPTED_CS_PDU | INFO_AUTOLOGON, + Domain: []byte{0, 0}, + UserName: []byte{0, 0}, + Password: []byte{0, 0}, + AlternateShell: []byte{0, 0}, + WorkingDir: []byte{0, 0}, + ExtendedInfo: NewExtendedInfo(nil), + } + // 批量压缩:声明 K64(64K 历史 MPPC)——core/mppc.go 解码器支持的变体。 + // 传统位图管线(16/24/32bpp 位图模式)不做此协商时服务器发原始位图, + // 带宽会高一个数量级。 + info.Flag |= INFO_COMPRESSION | INFO_CompressionTypeRDP5 + return info +} + +func (o *RDPInfo) SetClientAutoReconnect(auto *ClientAutoReconnect) { + o.ExtendedInfo.AutoReconnect = auto +} + +func (o *RDPInfo) SetClientInfo() { + o.Flag |= INFO_LOGONNOTIFY | INFO_LOGONERRORS +} + +func (o *RDPInfo) Serialize(hasExtended bool) []byte { + buff := &bytes.Buffer{} + core.WriteUInt32LE(o.CodePage, buff) // 0000000 + core.WriteUInt32LE(o.Flag, buff) // 0530101 + core.WriteUInt16LE(uint16(len(o.Domain)-2), buff) // 001c + core.WriteUInt16LE(uint16(len(o.UserName)-2), buff) // 0008 + core.WriteUInt16LE(uint16(len(o.Password)-2), buff) //000c + core.WriteUInt16LE(uint16(len(o.AlternateShell)-2), buff) //0000 + core.WriteUInt16LE(uint16(len(o.WorkingDir)-2), buff) //0000 + core.WriteBytes(o.Domain, buff) + core.WriteBytes(o.UserName, buff) + core.WriteBytes(o.Password, buff) + core.WriteBytes(o.AlternateShell, buff) + core.WriteBytes(o.WorkingDir, buff) + if hasExtended { + core.WriteBytes(o.ExtendedInfo.Serialize(), buff) + } + return buff.Bytes() +} + +type SecurityHeader struct { + securityFlag uint16 + securityFlagHi uint16 +} + +func readSecurityHeader(r io.Reader) *SecurityHeader { + s := &SecurityHeader{} + s.securityFlag, _ = core.ReadUint16LE(r) + s.securityFlagHi, _ = core.ReadUint16LE(r) + return s +} + +type SEC struct { + emission.Emitter + transport core.Transport + info *RDPInfo + machineName string + clientData []any + serverData []any + + enableEncryption bool + //Enable Secure Mac generation + enableSecureCheckSum bool + //counter before update + nbEncryptedPacket int + nbDecryptedPacket int + + currentDecrytKey []byte + currentEncryptKey []byte + + //current rc4 tab + decryptRc4 *rc4.Cipher + encryptRc4 *rc4.Cipher + + macKey []byte + + // fastPathSender is the underlying transport (typically TPKT) that knows + // how to wrap a payload in a fast-path frame. Set via SetFastPathSender + // to enable Fast-Path Client Input PDUs (MS-RDPBCGR §2.2.8.1.2). + fastPathSender core.FastPathSender +} + +func NewSEC(t core.Transport) *SEC { + sec := &SEC{ + *emission.NewEmitter(), + t, + NewRDPInfo(), + "", + nil, + nil, + false, + false, + 0, + 0, + nil, + nil, + nil, + nil, + nil, + nil, + } + + t.On("close", func() { + sec.Emit("close") + }).On("error", func(err error) { + sec.Emit("error", err) + }) + return sec +} + +func (s *SEC) Read(data []byte) (n int, err error) { + return s.transport.Read(data) +} + +func (s *SEC) Write(b []byte) (n int, err error) { + if !s.enableEncryption { + return s.transport.Write(b) + } + data := s.encrytData(b) + return s.transport.Write(data) +} + +func (s *SEC) Close() error { + return s.transport.Close() +} + +// SetFastPathSender wires the underlying transport that can frame fast-path +// PDUs. When set, SendFastPath is usable. +func (s *SEC) SetFastPathSender(f core.FastPathSender) { + s.fastPathSender = f +} + +// SendFastPath wraps the given payload in a fast-path frame using the +// underlying transport. Returns an error when legacy RDP encryption is +// enabled (this layer does not yet sign fast-path output), allowing the +// caller to fall back to a slow-path send. +func (s *SEC) SendFastPath(secFlag byte, b []byte) (int, error) { + if s.fastPathSender == nil { + return 0, fmt.Errorf("sec: fastPathSender not set") + } + if s.enableEncryption { + return 0, fmt.Errorf("sec: fast-path output not supported with legacy encryption") + } + return s.fastPathSender.SendFastPath(secFlag, b) +} + +// LegacyEncryptionEnabled reports whether the per-PDU RDP encryption layer +// (not TLS/CredSSP) is in use. Callers use this to disable optimisations +// like fast-path input that this layer does not yet implement signing for. +func (s *SEC) LegacyEncryptionEnabled() bool { + return s.enableEncryption +} + +func (s *SEC) sendFlagged(flag uint16, data []byte) (n int, err error) { + slog.Debug("sendFlagged", "flag", flag, "data", core.Hex(data)) + b := s.encryt(flag, data) + return s.transport.Write(b) +} + +/* +@see: http://msdn.microsoft.com/en-us/library/cc241995.aspx +@param macSaltKey: {str} mac key +@param data: {str} data to sign +@return: {str} signature +*/ +func macData(macSaltKey, data []byte) []byte { + sha1Digest := sha1.New() + md5Digest := md5.New() + + var lenBuf [4]byte + binary.LittleEndian.PutUint32(lenBuf[:], uint32(len(data))) + + sha1Digest.Write(macSaltKey) + sha1Digest.Write(macPad36[:]) + sha1Digest.Write(lenBuf[:]) + sha1Digest.Write(data) + + sha1Sig := sha1Digest.Sum(nil) + + md5Digest.Write(macSaltKey) + md5Digest.Write(macPad5C[:]) + md5Digest.Write(sha1Sig) + + return md5Digest.Sum(nil) +} +func (s *SEC) readEncryptedPayload(data []byte, checkSum bool) []byte { + sign := data[:8] + slog.Debug("readEncryptedPayload", "sign", sign) + encryptedPayload := data[8:] + if s.decryptRc4 == nil { + s.decryptRc4, _ = rc4.NewCipher(s.currentDecrytKey) + } + s.nbDecryptedPacket++ + plaintext := make([]byte, len(encryptedPayload)) + s.decryptRc4.XORKeyStream(plaintext, encryptedPayload) + + return plaintext +} +func (s *SEC) writeEncryptedPayload(data []byte, checkSum bool) []byte { + if checkSum { + return []byte{} + } + + s.nbEncryptedPacket++ + slog.Debug("writeEncryptedPayload", "nbEncryptedPacket", s.nbEncryptedPacket) + + sign := macData(s.macKey, data)[:8] + if s.encryptRc4 == nil { + s.encryptRc4, _ = rc4.NewCipher(s.currentEncryptKey) + } + + result := make([]byte, 8+len(data)) + copy(result[:8], sign) + s.encryptRc4.XORKeyStream(result[8:], data) + slog.Debug("writeEncryptedPayload", "sign", core.Hex(sign), "plaintext", core.Hex(result[8:])) + return result +} + +func (s *SEC) encryt(flag uint16, b []byte) []byte { + data := b + if flag&ENCRYPT != 0 { + data = s.writeEncryptedPayload(b, flag&SECURE_CHECKSUM != 0) + } + result := make([]byte, 4+len(data)) + binary.LittleEndian.PutUint16(result[0:], flag) + binary.LittleEndian.PutUint16(result[2:], 0) + copy(result[4:], data) + return result +} +func (s *SEC) encrytData(b []byte) []byte { + if !s.enableEncryption { + return b + } + + var flag uint16 = ENCRYPT + if s.enableSecureCheckSum { + flag |= SECURE_CHECKSUM + } + return s.encryt(flag, b) +} + +func (s *SEC) decrytData(b []byte) []byte { + if !s.enableEncryption { + return b + } + + if len(b) < 4 { + return b + } + securityFlag := binary.LittleEndian.Uint16(b[0:]) + // securityFlagHi = b[2:4] (ignored) + data := b[4:] + if securityFlag&ENCRYPT != 0 { + data = s.readEncryptedPayload(data, securityFlag&SECURE_CHECKSUM != 0) + } + return data +} + +type Client struct { + *SEC + userId uint16 + channelId uint16 + //initialise decrypt and encrypt keys + initialDecrytKey []byte + initialEncryptKey []byte + + fastPathListener core.FastPathListener + channelSender core.ChannelSender +} + +func NewClient(t core.Transport) *Client { + c := &Client{ + SEC: NewSEC(t), + } + t.On("connect", c.connect) + return c +} + +func (c *Client) SetClientAutoReconnect(id uint32, random []byte) { + auto := NewClientAutoReconnect(id, random) + c.info.SetClientAutoReconnect(auto) +} + +// SetPerformanceFlags 覆盖 Client Info PDU 的 performanceFlags +// (MS-RDPBCGR 2.2.1.11.1.1.1):禁用类位置位 = 关闭对应桌面元素, +// PERF_ENABLE_FONT_SMOOTHING / PERF_ENABLE_DESKTOP_COMPOSITION 置位 = +// 开启对应增强。未调用时保持 NewExtendedInfo 的默认(视觉全开)。 +func (c *Client) SetPerformanceFlags(flags uint32) { + c.info.ExtendedInfo.PerformanceFlags = flags +} + +// SetNoAudioPlayback 声明客户端不播放音频(mstsc「不播放」)。 +func (c *Client) SetNoAudioPlayback() { + c.info.Flag |= INFO_NOAUDIOPLAYBACK +} + +// SetRemoteConsoleAudio 声明音频在服务器本机播放(mstsc「在远程计算机播放」)。 +func (c *Client) SetRemoteConsoleAudio() { + c.info.Flag |= INFO_REMOTECONSOLEAUDIO +} + +func (c *Client) SetAlternateShell(shell string) { + buff := &bytes.Buffer{} + for _, ch := range utf16.Encode([]rune(shell)) { + core.WriteUInt16LE(ch, buff) + } + core.WriteUInt16LE(0, buff) + c.info.AlternateShell = buff.Bytes() + c.info.Flag |= INFO_RAIL +} + +func (c *Client) SetUser(user string) { + buff := &bytes.Buffer{} + for _, ch := range utf16.Encode([]rune(user)) { + core.WriteUInt16LE(ch, buff) + } + core.WriteUInt16LE(0, buff) + c.info.UserName = buff.Bytes() +} + +func (c *Client) SetPwd(pwd string) { + buff := &bytes.Buffer{} + for _, ch := range utf16.Encode([]rune(pwd)) { + core.WriteUInt16LE(ch, buff) + } + core.WriteUInt16LE(0, buff) + c.info.Password = buff.Bytes() +} + +func (c *Client) SetDomain(domain string) { + buff := &bytes.Buffer{} + for _, ch := range utf16.Encode([]rune(domain)) { + core.WriteUInt16LE(ch, buff) + } + core.WriteUInt16LE(0, buff) + c.info.Domain = buff.Bytes() +} + +// SetClientTimezone 按 [MS-RDPBCGR] 2.2.1.11.1.1.1.1 填充 Client Info PDU 时区。 +// name 为 Windows 时区注册表键名(如 "UTC"、"China Standard Time"); +// biasMinutes 为 UTC 与本地时间之差(东八区为 -480)。 +// dynamic DST 键名缺失或为空会使现代 Windows 服务器以 +// ERRINFO_TIMEZONE_KEY_NAME_LENGTH_TOO_SHORT (0x112F) 断开连接。 +func (c *Client) SetClientTimezone(name string, biasMinutes int) { + tz := make([]byte, 172) + binary.LittleEndian.PutUint32(tz[0:4], uint32(int32(biasMinutes))) + // 布局:bias(4) standardName(64) standardDate(16) standardBias(4) + // daylightName(64) daylightDate(16) daylightBias(4) + writeName := func(off int, s string) { + w := utf16.Encode([]rune(s)) + for i := 0; i*2+1 < 64; i++ { + v := uint16(0) + if i < len(w) { + v = w[i] + } + binary.LittleEndian.PutUint16(tz[off+i*2:], v) + } + } + writeName(4, name) // standardName + writeName(88, name) // daylightName + c.info.ExtendedInfo.ClientTimeZone = tz + + dyn := &bytes.Buffer{} + for _, ch := range utf16.Encode([]rune(name)) { + core.WriteUInt16LE(ch, dyn) + } + c.info.ExtendedInfo.DynDSTTimeZoneKeyName = dyn.Bytes() +} + +func (c *Client) connect(clientData []any, serverData []any, userId uint16, channels []t125.MCSChannelInfo) { + slog.Debug("connected!", "clientData", clientData, "serverData", serverData, "userId", userId, "channels", channels) + c.clientData = clientData + c.serverData = serverData + c.userId = userId + for _, channel := range channels { + if channel.Name == t125.GLOBAL_CHANNEL_NAME { + c.channelId = channel.ID + //break + } + } + c.enableEncryption = c.ClientCoreData().ServerSelectedProtocol == 0 + + if c.enableEncryption { + c.sendClientRandom() + } + + c.sendInfoPkt() + c.transport.Once("sec", c.recvLicenceInfo) +} + +func (c *Client) ClientCoreData() *gcc.ClientCoreData { + return c.clientData[0].(*gcc.ClientCoreData) +} +func (c *Client) ClientSecurityData() *gcc.ClientSecurityData { + return c.clientData[1].(*gcc.ClientSecurityData) +} +func (c *Client) ClientNetworkData() *gcc.ClientNetworkData { + return c.clientData[2].(*gcc.ClientNetworkData) +} + +func (c *Client) serverCoreData() *gcc.ServerCoreData { + return c.serverData[0].(*gcc.ServerCoreData) +} +func (c *Client) ServerSecurityData() *gcc.ServerSecurityData { + return c.serverData[1].(*gcc.ServerSecurityData) +} + +/* +@summary: generate 40 bits data from 128 bits data +@param data: {str} 128 bits data +@return: {str} 40 bits data +@see: http://msdn.microsoft.com/en-us/library/cc240785.aspx +*/ +func gen40bits(data []byte) []byte { + return append([]byte("\xd1\x26\x9e"), data[3:8]...) +} + +/* +@summary: generate 56 bits data from 128 bits data +@param data: {str} 128 bits data +@return: {str} 56 bits data +@see: http://msdn.microsoft.com/en-us/library/cc240785.aspx +*/ +func gen56bits(data []byte) []byte { + return append([]byte("\xd1"), data[1:8]...) +} + +/* +@summary: Generate particular signature from combination of sha1 and md5 +@see: http://msdn.microsoft.com/en-us/library/cc241992.aspx +@param inputData: strange input (see doc) +@param salt: salt for context call +@param salt1: another salt (ex : client random) +@param salt2: another another salt (ex: server random) +@return : MD5(Salt + SHA1(Input + Salt + Salt1 + Salt2)) +*/ +func saltedHash(inputData, salt, salt1, salt2 []byte) []byte { + sha1Digest := sha1.New() + md5Digest := md5.New() + + sha1Digest.Write(inputData) + sha1Digest.Write(salt[:48]) + sha1Digest.Write(salt1) + sha1Digest.Write(salt2) + sha1Sig := sha1Digest.Sum(nil) + + md5Digest.Write(salt[:48]) + md5Digest.Write(sha1Sig) + + return md5Digest.Sum(nil)[:16] +} + +/* +@summary: MD5(in0[:16] + in1[:32] + in2[:32]) +@param key: in 16 +@param random1: in 32 +@param random2: in 32 +@return MD5(in0[:16] + in1[:32] + in2[:32]) +*/ +func finalHash(key, random1, random2 []byte) []byte { + md5Digest := md5.New() + md5Digest.Write(key) + md5Digest.Write(random1) + md5Digest.Write(random2) + return md5Digest.Sum(nil) +} + +/* +@summary: Generate master secret +@param secret: {str} secret +@param clientRandom : {str} client random +@param serverRandom : {str} server random +@see: http://msdn.microsoft.com/en-us/library/cc241992.aspx +*/ +func masterSecret(secret, random1, random2 []byte) []byte { + sh1 := saltedHash([]byte("A"), secret, random1, random2) + sh2 := saltedHash([]byte("BB"), secret, random1, random2) + sh3 := saltedHash([]byte("CCC"), secret, random1, random2) + ms := bytes.NewBuffer(nil) + ms.Write(sh1) + ms.Write(sh2) + ms.Write(sh3) + return ms.Bytes() +} + +/* +@summary: Generate master secret +@param secret: secret +@param clientRandom : client random +@param serverRandom : server random +*/ +func sessionKeyBlob(secret, random1, random2 []byte) []byte { + sh1 := saltedHash([]byte("X"), secret, random1, random2) + sh2 := saltedHash([]byte("YY"), secret, random1, random2) + sh3 := saltedHash([]byte("ZZZ"), secret, random1, random2) + ms := bytes.NewBuffer(nil) + ms.Write(sh1) + ms.Write(sh2) + ms.Write(sh3) + return ms.Bytes() + +} +func generateKeys(clientRandom, serverRandom []byte, method uint32) ([]byte, []byte, []byte) { + b := &bytes.Buffer{} + b.Write(clientRandom[:24]) + b.Write(serverRandom[:24]) + preMasterHash := b.Bytes() + slog.Debug("getnerateKeys", "method", method) + + masterHash := masterSecret(preMasterHash, clientRandom, serverRandom) + sessionKey := sessionKeyBlob(masterHash, clientRandom, serverRandom) + macKey128 := sessionKey[:16] + initialFirstKey128 := finalHash(sessionKey[16:32], clientRandom, serverRandom) + initialSecondKey128 := finalHash(sessionKey[32:48], clientRandom, serverRandom) + + //generate valid key + if method == gcc.ENCRYPTION_FLAG_40BIT { + return gen40bits(macKey128), gen40bits(initialFirstKey128), gen40bits(initialSecondKey128) + } else if method == gcc.ENCRYPTION_FLAG_56BIT { + return gen56bits(macKey128), gen56bits(initialFirstKey128), gen56bits(initialSecondKey128) + } + // method == gcc.ENCRYPTION_FLAG_128BIT + return macKey128, initialFirstKey128, initialSecondKey128 + +} + +type ClientSecurityExchangePDU struct { + Length uint32 `struc:"little"` + EncryptedClientRandom []byte `struc:"little"` + Padding []byte `struc:"[8]byte"` +} + +func (e *ClientSecurityExchangePDU) serialize() []byte { + buff := &bytes.Buffer{} + core.WriteUInt32LE(e.Length, buff) + core.WriteBytes(e.EncryptedClientRandom, buff) + core.WriteBytes(e.Padding, buff) + + return buff.Bytes() +} +func (c *Client) sendClientRandom() { + clientRandom := core.Random(32) + slog.Debug("sendClientRandom", "clientRandom", core.Hex(clientRandom)) + + serverRandom := c.ServerSecurityData().ServerRandom + slog.Debug("sendlientRandom", "ServerRandom", core.Hex(serverRandom)) + + c.macKey, c.initialDecrytKey, c.initialEncryptKey = generateKeys(clientRandom, + serverRandom, c.ServerSecurityData().EncryptionMethod) + + //initialize keys + c.currentDecrytKey = c.initialDecrytKey + c.currentEncryptKey = c.initialEncryptKey + + //verify certificate + if !c.ServerSecurityData().ServerCertificate.CertData.Verify() { + slog.Warn("Cannot verify server identity") + } + + serverPubKey, _ := c.ServerSecurityData().ServerCertificate.CertData.GetPublicKey() + ret, err := rsa.EncryptPKCS1v15(rand.Reader, serverPubKey, core.Reverse(clientRandom)) + if err != nil { + slog.Error("sendlientRandom", "err", err) + } + message := ClientSecurityExchangePDU{} + message.EncryptedClientRandom = core.Reverse(ret) + message.Length = uint32(len(message.EncryptedClientRandom) + 8) + message.Padding = make([]byte, 8) + + slog.Debug("sendlientRandom", "message", message) + + c.sendFlagged(EXCHANGE_PKT, message.serialize()) +} +func (c *Client) sendInfoPkt() { + var secFlag uint16 = INFO_PKT + if c.enableEncryption { + secFlag |= ENCRYPT + } + + slog.Debug("sendInfoPkt", "secFlag", secFlag, "hasExtended", c.ClientCoreData().RdpVersion >= gcc.RDP_VERSION_5_PLUS, + "infoFlag", fmt.Sprintf("0x%08X", c.info.Flag)) + infoBytes := c.info.Serialize(c.ClientCoreData().RdpVersion >= gcc.RDP_VERSION_5_PLUS) + slog.Debug("sendInfoPkt bytes", "len", len(infoBytes), "hex", core.Hex(infoBytes)) + c.sendFlagged(secFlag, infoBytes) +} + +func (c *Client) recvLicenceInfo(channel string, s []byte) { + slog.Debug("recvLicenceInfo", "s", core.Hex(s)) + r := bytes.NewReader(s) + h := readSecurityHeader(r) + if (h.securityFlag & LICENSE_PKT) == 0 { + c.Emit("error", errors.New("NODE_RDP_PROTOCOL_PDU_SEC_BAD_LICENSE_HEADER")) + return + } + + p := lic.ReadLicensePacket(r) + switch p.BMsgtype { + case lic.NEW_LICENSE: + slog.Debug("sec NEW_LICENSE") + c.Emit("success") + goto connect + case lic.ERROR_ALERT: + message := p.LicensingMessage.(*lic.ErrorMessage) + slog.Debug("recvLicenceInfo ERROR_ALERT", "ErrorCode", message.DwErrorCode) + if message.DwErrorCode == lic.STATUS_VALID_CLIENT && message.DwStateTransaction == lic.ST_NO_TRANSITION { + goto connect + } + goto retry + case lic.LICENSE_REQUEST: + slog.Debug("recvLicenceInfo LICENSE_REQUEST") + c.sendClientNewLicenseRequest(p.LicensingMessage.([]byte)) + goto retry + case lic.PLATFORM_CHALLENGE: + slog.Debug("recvLicenceInfo PLATFORM_CHALLENGE") + c.sendClientChallengeResponse(p.LicensingMessage.([]byte)) + goto retry + default: + slog.Error("Not a valid license packet") + c.Emit("error", errors.New("Not a valid license packet")) + return + } + +connect: + c.transport.On("sec", c.recvData) + c.Emit("connect", c.clientData[0].(*gcc.ClientCoreData), c.userId, c.channelId) + return + +retry: + c.transport.Once("sec", c.recvLicenceInfo) + return +} + +func (c *Client) sendClientNewLicenseRequest(data []byte) { + var req lic.ServerLicenseRequest + struc.Unpack(bytes.NewReader(data), &req) + + var sc gcc.ServerCertificate + if c.ServerSecurityData().ServerCertificate.DwVersion != 0 { + sc = c.ServerSecurityData().ServerCertificate + } else { + rd := bytes.NewReader(req.ServerCertificate.BlobData) + err := sc.Unpack(rd) + if err != nil { + slog.Error("sendClientNewLicenseRequest", "err", err) + return + } + } + + serverRandom := req.ServerRandom + clientRandom := core.Random(32) + preMasterSecret := core.Random(48) + masSecret := masterSecret(preMasterSecret, clientRandom, serverRandom) + sessionKeyBlob := masterSecret(masSecret, serverRandom, clientRandom) + c.macKey = sessionKeyBlob[:16] + c.initialDecrytKey = finalHash(sessionKeyBlob[16:32], clientRandom, serverRandom) + + //format message + message := &lic.ClientNewLicenseRequest{} + message.PreferredKeyExchangeAlg = 0x00000001 + message.PlatformId = 0x04000000 | 0x00010000 + message.ClientRandom = clientRandom + + buff := &bytes.Buffer{} + + serverPubKey, _ := sc.CertData.GetPublicKey() + ret, err := rsa.EncryptPKCS1v15(rand.Reader, serverPubKey, core.Reverse(preMasterSecret)) + if err != nil { + slog.Error("sendClientNewLicenseRequest", "err", err) + } + + buff.Write(core.Reverse(ret)) + buff.Write([]byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}) + message.EncryptedPreMasterSecret.BlobData = buff.Bytes() + message.EncryptedPreMasterSecret.WBlobLen = uint16(buff.Len()) + message.EncryptedPreMasterSecret.WBlobType = lic.BB_RANDOM_BLOB + + buff.Reset() + buff.Write(c.info.UserName) + buff.Write([]byte{0x00}) + message.ClientUserName.BlobData = buff.Bytes() + message.ClientUserName.WBlobLen = uint16(buff.Len()) + message.ClientUserName.WBlobType = lic.BB_CLIENT_USER_NAME_BLOB + + buff.Reset() + buff.Write(c.ClientCoreData().ClientName[:]) + buff.Write([]byte{0x00}) + message.ClientMachineName.BlobData = buff.Bytes() + message.ClientMachineName.WBlobLen = uint16(buff.Len()) + message.ClientMachineName.WBlobType = lic.BB_CLIENT_MACHINE_NAME_BLOB + + buff.Reset() + err = struc.Pack(buff, message) + if err != nil { + slog.Error("sendClientNewLicenseRequest", "err", err) + } + + c.sendFlagged(LICENSE_PKT, buff.Bytes()) +} + +func (c *Client) sendClientChallengeResponse(data []byte) { + var pc lic.ServerPlatformChallenge + struc.Unpack(bytes.NewReader(data), &pc) + + serverEncryptedChallenge := pc.EncryptedPlatformChallenge.BlobData + //decrypt server challenge + //it should be TEST word in unicode format + rc, _ := rc4.NewCipher(c.initialDecrytKey) + serverChallenge := make([]byte, 20) + rc.XORKeyStream(serverChallenge, serverEncryptedChallenge) + //if serverChallenge != "T\x00E\x00S\x00T\x00\x00\x00": + //raise InvalidExpectedDataException("bad license server challenge") + + //generate hwid + b := &bytes.Buffer{} + b.Write(c.ClientCoreData().ClientName[:]) + b.Write(c.info.UserName) + for range 2 { + b.Write([]byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}) + } + hwid := b.Bytes()[:20] + + encryptedHWID := make([]byte, 20) + rc.XORKeyStream(encryptedHWID, hwid) + + b.Reset() + b.Write(serverChallenge) + b.Write(hwid) + + message := &lic.ClientPLatformChallengeResponse{} + message.EncryptedPlatformChallengeResponse.BlobData = serverEncryptedChallenge + message.EncryptedHWID.BlobData = encryptedHWID + message.MACData = macData(c.macKey, b.Bytes())[:16] + + b.Reset() + struc.Pack(b, message) + c.sendFlagged(LICENSE_PKT, b.Bytes()) +} + +func (c *Client) recvData(channel string, s []byte) { + data := c.decrytData(s) + if channel != t125.GLOBAL_CHANNEL_NAME { + c.Emit("channel", channel, data) + return + } + c.Emit("data", data) +} +func (c *Client) SetFastPathListener(f core.FastPathListener) { + c.fastPathListener = f +} + +func (c *Client) RecvFastPath(secFlag byte, s []byte) { + data := s + if c.enableEncryption && secFlag&FASTPATH_OUTPUT_ENCRYPTED != 0 { + data = c.readEncryptedPayload(s, secFlag&FASTPATH_OUTPUT_SECURE_CHECKSUM != 0) + } + c.fastPathListener.RecvFastPath(secFlag, data) +} + +func (c *Client) SetChannelSender(f core.ChannelSender) { + c.channelSender = f +} + +func (c *Client) SendToChannel(channel string, b []byte) (int, error) { + if !c.enableEncryption { + return c.channelSender.SendToChannel(channel, b) + } + var flag uint16 = ENCRYPT + if c.enableSecureCheckSum { + flag |= SECURE_CHECKSUM + } + data := c.writeEncryptedPayload(b, c.enableSecureCheckSum) + + buff := &bytes.Buffer{} + core.WriteUInt16LE(flag, buff) + core.WriteUInt16LE(0, buff) + core.WriteBytes(data, buff) + return c.channelSender.SendToChannel(channel, buff.Bytes()) +} diff --git a/protocol/t125/ber/ber.go b/protocol/t125/ber/ber.go new file mode 100644 index 0000000..280c17c --- /dev/null +++ b/protocol/t125/ber/ber.go @@ -0,0 +1,189 @@ +package ber + +import ( + "errors" + "fmt" + "io" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +const ( + CLASS_MASK uint8 = 0xC0 + CLASS_UNIV = 0x00 + CLASS_APPL = 0x40 + CLASS_CTXT = 0x80 + CLASS_PRIV = 0xC0 +) + +const ( + PC_MASK uint8 = 0x20 + PC_PRIMITIVE = 0x00 + PC_CONSTRUCT = 0x20 +) + +const ( + TAG_MASK uint8 = 0x1F + TAG_BOOLEAN = 0x01 + TAG_INTEGER = 0x02 + TAG_BIT_STRING = 0x03 + TAG_OCTET_STRING = 0x04 + TAG_OBJECT_IDENFIER = 0x06 + TAG_ENUMERATED = 0x0A + TAG_SEQUENCE = 0x10 + TAG_SEQUENCE_OF = 0x10 +) + +func berPC(pc bool) uint8 { + if pc { + return PC_CONSTRUCT + } + return PC_PRIMITIVE +} + +func ReadEnumerated(r io.Reader) (uint8, error) { + if !ReadUniversalTag(TAG_ENUMERATED, false, r) { + return 0, errors.New("invalid ber tag") + } + length, err := ReadLength(r) + if err != nil { + return 0, err + } + if length != 1 { + return 0, errors.New(fmt.Sprintf("enumerate size is wrong, get %v, expect 1", length)) + } + return core.ReadUInt8(r) +} + +func ReadUniversalTag(tag uint8, pc bool, r io.Reader) bool { + bb, _ := core.ReadUInt8(r) + return bb == (CLASS_UNIV|berPC(pc))|(TAG_MASK&tag) +} + +func WriteUniversalTag(tag uint8, pc bool, w io.Writer) { + core.WriteUInt8((CLASS_UNIV|berPC(pc))|(TAG_MASK&tag), w) +} + +func ReadLength(r io.Reader) (int, error) { + ret := 0 + size, _ := core.ReadUInt8(r) + if size&0x80 > 0 { + size = size &^ 0x80 + if size == 1 { + r, err := core.ReadUInt8(r) + if err != nil { + return 0, err + } + ret = int(r) + } else if size == 2 { + r, err := core.ReadUint16BE(r) + if err != nil { + return 0, err + } + ret = int(r) + } else { + return 0, errors.New("BER length may be 1 or 2") + } + } else { + ret = int(size) + } + return ret, nil +} + +func WriteLength(size int, w io.Writer) { + if size > 0x7f { + core.WriteUInt8(0x82, w) + core.WriteUInt16BE(uint16(size), w) + } else { + core.WriteUInt8(uint8(size), w) + } +} + +func ReadInteger(r io.Reader) (int, error) { + if !ReadUniversalTag(TAG_INTEGER, false, r) { + return 0, errors.New("Bad integer tag") + } + size, _ := ReadLength(r) + switch size { + case 1: + num, _ := core.ReadUInt8(r) + return int(num), nil + case 2: + num, _ := core.ReadUint16BE(r) + return int(num), nil + case 3: + integer1, _ := core.ReadUInt8(r) + integer2, _ := core.ReadUint16BE(r) + return int(integer2) + int(uint32(integer1)<<16), nil + case 4: + num, _ := core.ReadUInt32BE(r) + return int(num), nil + default: + return 0, errors.New("wrong size") + } +} + +func WriteInteger(n int, w io.Writer) { + WriteUniversalTag(TAG_INTEGER, false, w) + if n <= 0xff { + WriteLength(1, w) + core.WriteUInt8(uint8(n), w) + } else if n <= 0xffff { + WriteLength(2, w) + core.WriteUInt16BE(uint16(n), w) + } else { + WriteLength(4, w) + core.WriteUInt32BE(uint32(n), w) + } +} + +func WriteOctetstring(str string, w io.Writer) { + WriteUniversalTag(TAG_OCTET_STRING, false, w) + WriteLength(len(str), w) + core.WriteBytes([]byte(str), w) +} + +func WriteBoolean(b bool, w io.Writer) { + bb := uint8(0) + if b { + bb = uint8(0xff) + } + WriteUniversalTag(TAG_BOOLEAN, false, w) + WriteLength(1, w) + core.WriteUInt8(bb, w) +} + +func ReadApplicationTag(tag uint8, r io.Reader) (int, error) { + bb, _ := core.ReadUInt8(r) + if tag > 30 { + if bb != (CLASS_APPL|PC_CONSTRUCT)|TAG_MASK { + return 0, errors.New(fmt.Sprintf("ReadApplicationTag error tag=0x%x,bb=0x%x", tag, bb)) + } + bb, _ := core.ReadUInt8(r) + if bb != tag { + return 0, errors.New("ReadApplicationTag bad tag") + } + } else { + if bb != (CLASS_APPL|PC_CONSTRUCT)|(TAG_MASK&tag) { + return 0, errors.New(fmt.Sprintf("ReadApplicationTag error valiable length tag=0x%x,bb=0x%x", tag, bb)) + } + } + return ReadLength(r) +} + +func WriteApplicationTag(tag uint8, size int, w io.Writer) { + if tag > 30 { + core.WriteUInt8((CLASS_APPL|PC_CONSTRUCT)|TAG_MASK, w) + core.WriteUInt8(tag, w) + WriteLength(size, w) + } else { + core.WriteUInt8((CLASS_APPL|PC_CONSTRUCT)|(TAG_MASK&tag), w) + WriteLength(size, w) + } +} + +func WriteEncodedDomainParams(data []byte, w io.Writer) { + WriteUniversalTag(TAG_SEQUENCE, true, w) + WriteLength(len(data), w) + core.WriteBytes(data, w) +} diff --git a/protocol/t125/gcc/gcc.go b/protocol/t125/gcc/gcc.go new file mode 100644 index 0000000..407e298 --- /dev/null +++ b/protocol/t125/gcc/gcc.go @@ -0,0 +1,671 @@ +package gcc + +import ( + "bytes" + "crypto/rsa" + "crypto/x509" + "crypto/x509/pkix" + "encoding/asn1" + "errors" + "io" + "log/slog" + "math/big" + "os" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/per" +) + +var t124_02_98_oid = []byte{0, 0, 20, 124, 0, 1} +var h221_cs_key = "Duca" +var h221_sc_key = "McDn" + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240509.aspx + */ +type Message uint16 + +const ( + //server -> client + SC_CORE Message = 0x0C01 + SC_SECURITY = 0x0C02 + SC_NET = 0x0C03 + SC_MCS_MSGCHANNEL = 0x0C04 // TS_UD_SC_MCS_MSGCHANNEL + //client -> server + CS_CORE = 0xC001 + CS_SECURITY = 0xC002 + CS_NET = 0xC003 + CS_CLUSTER = 0xC004 + CS_MONITOR = 0xC005 + CS_MCS_MSGCHANNEL = 0xC006 // TS_UD_CS_MCS_MSGCHANNEL +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type ColorDepth uint16 + +const ( + RNS_UD_COLOR_8BPP ColorDepth = 0xCA01 + RNS_UD_COLOR_16BPP_555 = 0xCA02 + RNS_UD_COLOR_16BPP_565 = 0xCA03 + RNS_UD_COLOR_24BPP = 0xCA04 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type HighColor uint16 + +const ( + HIGH_COLOR_4BPP HighColor = 0x0004 + HIGH_COLOR_8BPP = 0x0008 + HIGH_COLOR_15BPP = 0x000f + HIGH_COLOR_16BPP = 0x0010 + HIGH_COLOR_24BPP = 0x0018 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type Support uint16 + +const ( + RNS_UD_24BPP_SUPPORT uint16 = 0x0001 + RNS_UD_16BPP_SUPPORT = 0x0002 + RNS_UD_15BPP_SUPPORT = 0x0004 + RNS_UD_32BPP_SUPPORT = 0x0008 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type CapabilityFlag uint16 + +const ( + RNS_UD_CS_SUPPORT_ERRINFO_PDU uint16 = 0x0001 + RNS_UD_CS_WANT_32BPP_SESSION = 0x0002 + RNS_UD_CS_SUPPORT_STATUSINFO_PDU = 0x0004 + RNS_UD_CS_STRONG_ASYMMETRIC_KEYS = 0x0008 + RNS_UD_CS_UNUSED = 0x0010 + RNS_UD_CS_VALID_CONNECTION_TYPE = 0x0020 + RNS_UD_CS_SUPPORT_MONITOR_LAYOUT_PDU = 0x0040 + RNS_UD_CS_SUPPORT_NETCHAR_AUTODETECT = 0x0080 + RNS_UD_CS_SUPPORT_DYNVC_GFX_PROTOCOL = 0x0100 + RNS_UD_CS_SUPPORT_DYNAMIC_TIME_ZONE = 0x0200 + RNS_UD_CS_SUPPORT_HEARTBEAT_PDU = 0x0400 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type ConnectionType uint8 + +const ( + CONNECTION_TYPE_MODEM ConnectionType = 0x01 + CONNECTION_TYPE_BROADBAND_LOW = 0x02 + CONNECTION_TYPE_SATELLITEV = 0x03 + CONNECTION_TYPE_BROADBAND_HIGH = 0x04 + CONNECTION_TYPE_WAN = 0x05 + CONNECTION_TYPE_LAN = 0x06 + CONNECTION_TYPE_AUTODETECT = 0x07 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240510.aspx + */ +type VERSION uint32 + +const ( + RDP_VERSION_4 VERSION = 0x00080001 + RDP_VERSION_5_PLUS = 0x00080004 + RDP_VERSION_10 = 0x00080005 + RDP_VERSION_10_1 = 0x00080006 + RDP_VERSION_10_2 = 0x00080007 +) + +type Sequence uint16 + +const ( + RNS_UD_SAS_DEL Sequence = 0xAA03 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240511.aspx + */ +type EncryptionMethod uint32 + +const ( + ENCRYPTION_FLAG_40BIT uint32 = 0x00000001 + ENCRYPTION_FLAG_128BIT = 0x00000002 + ENCRYPTION_FLAG_56BIT = 0x00000008 + FIPS_ENCRYPTION_FLAG = 0x00000010 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240518.aspx + */ +type EncryptionLevel uint32 + +const ( + ENCRYPTION_LEVEL_NONE EncryptionLevel = 0x00000000 + ENCRYPTION_LEVEL_LOW = 0x00000001 + ENCRYPTION_LEVEL_CLIENT_COMPATIBLE = 0x00000002 + ENCRYPTION_LEVEL_HIGH = 0x00000003 + ENCRYPTION_LEVEL_FIPS = 0x00000004 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240513.aspx + */ +type ChannelOptions uint32 + +const ( + CHANNEL_OPTION_INITIALIZED ChannelOptions = 0x80000000 + CHANNEL_OPTION_ENCRYPT_RDP = 0x40000000 + CHANNEL_OPTION_ENCRYPT_SC = 0x20000000 + CHANNEL_OPTION_ENCRYPT_CS = 0x10000000 + CHANNEL_OPTION_PRI_HIGH = 0x08000000 + CHANNEL_OPTION_PRI_MED = 0x04000000 + CHANNEL_OPTION_PRI_LOW = 0x02000000 + CHANNEL_OPTION_COMPRESS_RDP = 0x00800000 + CHANNEL_OPTION_COMPRESS = 0x00400000 + CHANNEL_OPTION_SHOW_PROTOCOL = 0x00200000 + REMOTE_CONTROL_PERSISTENT = 0x00100000 +) + +/** + * IBM_101_102_KEYS is the most common keyboard type + */ +type KeyboardType uint32 + +const ( + KT_IBM_PC_XT_83_KEY KeyboardType = 0x00000001 + KT_OLIVETTI = 0x00000002 + KT_IBM_PC_AT_84_KEY = 0x00000003 + KT_IBM_101_102_KEYS = 0x00000004 + KT_NOKIA_1050 = 0x00000005 + KT_NOKIA_9140 = 0x00000006 + KT_JAPANESE = 0x00000007 +) + +/** + * @see http://technet.microsoft.com/en-us/library/cc766503%28WS.10%29.aspx + */ +type KeyboardLayout uint32 + +const ( + ARABIC KeyboardLayout = 0x00000401 + BULGARIAN = 0x00000402 + CHINESE_US_KEYBOARD = 0x00000404 + CZECH = 0x00000405 + DANISH = 0x00000406 + GERMAN = 0x00000407 + GREEK = 0x00000408 + US = 0x00000409 + SPANISH = 0x0000040a + FINNISH = 0x0000040b + FRENCH = 0x0000040c + HEBREW = 0x0000040d + HUNGARIAN = 0x0000040e + ICELANDIC = 0x0000040f + ITALIAN = 0x00000410 + JAPANESE = 0x00000411 + KOREAN = 0x00000412 + DUTCH = 0x00000413 + NORWEGIAN = 0x00000414 +) + +/** + * @see http://msdn.microsoft.com/en-us/library/cc240521.aspx + */ +type CertificateType uint32 + +const ( + CERT_CHAIN_VERSION_1 CertificateType = 0x00000001 + CERT_CHAIN_VERSION_2 = 0x00000002 +) + +type ChannelDef struct { + Name string `struc:"little"` + Options uint32 `struc:"little"` +} + +type ClientCoreData struct { + RdpVersion VERSION `struc:"uint32,little"` + DesktopWidth uint16 `struc:"little"` + DesktopHeight uint16 `struc:"little"` + ColorDepth ColorDepth `struc:"little"` + SasSequence Sequence `struc:"little"` + KbdLayout KeyboardLayout `struc:"little"` + ClientBuild uint32 `struc:"little"` + ClientName [32]byte `struc:"[32]byte"` + KeyboardType uint32 `struc:"little"` + KeyboardSubType uint32 `struc:"little"` + KeyboardFnKeys uint32 `struc:"little"` + ImeFileName [64]byte `struc:"[64]byte"` + PostBeta2ColorDepth ColorDepth `struc:"little"` + ClientProductId uint16 `struc:"little"` + SerialNumber uint32 `struc:"little"` + HighColorDepth HighColor `struc:"little"` + SupportedColorDepths uint16 `struc:"little"` + EarlyCapabilityFlags uint16 `struc:"little"` + ClientDigProductId [64]byte `struc:"[64]byte"` + ConnectionType uint8 `struc:"uint8"` + Pad1octet uint8 `struc:"uint8"` + ServerSelectedProtocol uint32 `struc:"little"` +} + +func NewClientCoreData(kbdLayout uint32, keyboardType uint32, keyboardSubType uint32) *ClientCoreData { + name, _ := os.Hostname() + var ClientName [32]byte + copy(ClientName[:], core.UnicodeEncode(name)[:]) + return &ClientCoreData{ + RDP_VERSION_10_2, 1280, 800, RNS_UD_COLOR_8BPP, + RNS_UD_SAS_DEL, KeyboardLayout(kbdLayout), 22621, ClientName, keyboardType, + keyboardSubType, 12, [64]byte{}, RNS_UD_COLOR_8BPP, 1, 0, HIGH_COLOR_24BPP, + RNS_UD_15BPP_SUPPORT | RNS_UD_16BPP_SUPPORT | RNS_UD_24BPP_SUPPORT | RNS_UD_32BPP_SUPPORT, + RNS_UD_CS_SUPPORT_ERRINFO_PDU | RNS_UD_CS_WANT_32BPP_SESSION | RNS_UD_CS_VALID_CONNECTION_TYPE | RNS_UD_CS_SUPPORT_NETCHAR_AUTODETECT, + [64]byte{}, uint8(CONNECTION_TYPE_LAN), 0, 0} +} + +func (data *ClientCoreData) Pack() []byte { + buff := &bytes.Buffer{} + core.WriteUInt16LE(CS_CORE, buff) // 01C0 + core.WriteUInt16LE(0xd8, buff) // d800 + struc.Pack(buff, data) + return buff.Bytes() +} + +type ClientNetworkData struct { + ChannelCount uint32 + ChannelDefArray []ChannelDef +} + +func NewClientNetworkData() *ClientNetworkData { + n := &ClientNetworkData{ChannelDefArray: make([]ChannelDef, 0, 100)} + + /*var d1 ChannelDef + d1.Name = plugin.RDPDR_SVC_CHANNEL_NAME + d1.Options = uint32(CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | + CHANNEL_OPTION_COMPRESS_RDP) + n.ChannelDefArray = append(n.ChannelDefArray, d1) + + var d2 ChannelDef + d2.Name = plugin.RDPSND_SVC_CHANNEL_NAME + d2.Options = uint32(CHANNEL_OPTION_INITIALIZED | CHANNEL_OPTION_ENCRYPT_RDP | + CHANNEL_OPTION_COMPRESS_RDP | CHANNEL_OPTION_SHOW_PROTOCOL) + n.ChannelDefArray = append(n.ChannelDefArray, d2)*/ + + return n +} + +func (n *ClientNetworkData) AddVirtualChannel(name string, option uint32) { + var d ChannelDef + d.Name = name + d.Options = option + n.ChannelDefArray = append(n.ChannelDefArray, d) + n.ChannelCount++ +} + +func (n *ClientNetworkData) Pack() []byte { + buff := &bytes.Buffer{} + core.WriteUInt16LE(CS_NET, buff) // type + length := uint16(n.ChannelCount*12 + 8) + core.WriteUInt16LE(length, buff) // len 8 + core.WriteUInt32LE(n.ChannelCount, buff) + for i := 0; i < int(n.ChannelCount); i++ { + v := n.ChannelDefArray[i] + var name [8]byte + copy(name[:], v.Name) + core.WriteBytes(name[:], buff) + core.WriteUInt32LE(v.Options, buff) + } + return buff.Bytes() +} + +type ClientSecurityData struct { + EncryptionMethods uint32 + ExtEncryptionMethods uint32 +} + +func NewClientSecurityData() *ClientSecurityData { + return &ClientSecurityData{ + ENCRYPTION_FLAG_40BIT | ENCRYPTION_FLAG_56BIT | ENCRYPTION_FLAG_128BIT, + 00} +} + +func (d *ClientSecurityData) Pack() []byte { + buff := &bytes.Buffer{} + core.WriteUInt16LE(CS_SECURITY, buff) // type + core.WriteUInt16LE(0x0c, buff) // len 12 + core.WriteUInt32LE(d.EncryptionMethods, buff) + core.WriteUInt32LE(d.ExtEncryptionMethods, buff) + return buff.Bytes() +} + +type RSAPublicKey struct { + Magic uint32 `struc:"little"` //0x31415352 + Keylen uint32 `struc:"little,sizeof=Modulus"` + Bitlen uint32 `struc:"little"` + Datalen uint32 `struc:"little"` + PubExp uint32 `struc:"little"` + Modulus []byte `struc:"little"` + Padding []byte `struc:"[8]byte"` +} + +type ProprietaryServerCertificate struct { + DwSigAlgId uint32 `struc:"little"` //0x00000001 + DwKeyAlgId uint32 `struc:"little"` //0x00000001 + PublicKeyBlobType uint16 `struc:"little"` //0x0006 + PublicKeyBlobLen uint16 `struc:"little,sizeof=PublicKeyBlob"` + PublicKeyBlob RSAPublicKey `struc:"little"` + SignatureBlobType uint16 `struc:"little"` //0x0008 + SignatureBlobLen uint16 `struc:"little,sizeof=SignatureBlob"` + SignatureBlob []byte `struc:"little"` + //PaddingLen uint16 `struc:"little,sizeof=Padding,skip"` + Padding []byte `struc:"[8]byte"` +} + +func (p *ProprietaryServerCertificate) GetPublicKey() (*rsa.PublicKey, error) { + b := new(big.Int).SetBytes(core.Reverse(p.PublicKeyBlob.Modulus)) + e := new(big.Int).SetInt64(int64(p.PublicKeyBlob.PubExp)) + return &rsa.PublicKey{N: b, E: int(e.Int64())}, nil +} +func (p *ProprietaryServerCertificate) Verify() bool { + return true +} +func (p *ProprietaryServerCertificate) Encrypt() []byte { + //todo + return nil +} +func (p *ProprietaryServerCertificate) Unpack(r io.Reader) error { + p.DwSigAlgId, _ = core.ReadUInt32LE(r) + p.DwKeyAlgId, _ = core.ReadUInt32LE(r) + p.PublicKeyBlobType, _ = core.ReadUint16LE(r) + p.PublicKeyBlobLen, _ = core.ReadUint16LE(r) + var b RSAPublicKey + b.Magic, _ = core.ReadUInt32LE(r) + b.Keylen, _ = core.ReadUInt32LE(r) + b.Bitlen, _ = core.ReadUInt32LE(r) + b.Datalen, _ = core.ReadUInt32LE(r) + b.PubExp, _ = core.ReadUInt32LE(r) + b.Modulus, _ = core.ReadBytes(int(b.Keylen)-8, r) + b.Padding, _ = core.ReadBytes(8, r) + p.PublicKeyBlob = b + p.SignatureBlobType, _ = core.ReadUint16LE(r) + p.SignatureBlobLen, _ = core.ReadUint16LE(r) + p.SignatureBlob, _ = core.ReadBytes(int(p.SignatureBlobLen)-8, r) + p.Padding, _ = core.ReadBytes(8, r) + + return nil +} + +type CertBlob struct { + CbCert uint32 `struc:"little,sizeof=AbCert"` + AbCert []byte `struc:"little"` +} +type X509CertificateChain struct { + NumCertBlobs uint32 `struc:"little,sizeof=CertBlobArray"` + CertBlobArray []CertBlob `struc:"little"` + Padding []byte `struc:"[12]byte"` +} + +func (x *X509CertificateChain) GetPublicKey() (*rsa.PublicKey, error) { + data := x.CertBlobArray[len(x.CertBlobArray)-1].AbCert + cert, err := x509.ParseCertificate(data) + if err != nil { + slog.Error("X509 ParseCertificate", "err", err) + return nil, err + } + var rsaPublicKey *rsa.PublicKey + if cert.PublicKey == nil { + var pubKeyInfo struct { + Algorithm pkix.AlgorithmIdentifier + SubjectPublicKey asn1.BitString + } + _, err = asn1.Unmarshal(cert.RawSubjectPublicKeyInfo, &pubKeyInfo) + if err != nil { + return nil, err + } + rsaPublicKey, err = x509.ParsePKCS1PublicKey(pubKeyInfo.SubjectPublicKey.Bytes) + if err != nil { + return nil, err + } + } else { + rsaPublicKey = cert.PublicKey.(*rsa.PublicKey) + } + + return rsaPublicKey, nil +} +func (x *X509CertificateChain) Verify() bool { + return true +} +func (x *X509CertificateChain) Encrypt() []byte { + //todo + return nil +} +func (x *X509CertificateChain) Unpack(r io.Reader) error { + return struc.Unpack(r, x) +} + +type ServerCoreData struct { + RdpVersion VERSION `struc:"uint32,little"` + ClientRequestedProtocol uint32 `struc:"little"` + EarlyCapabilityFlags uint32 `struc:"little"` +} + +func NewServerCoreData() *ServerCoreData { + return &ServerCoreData{ + RDP_VERSION_5_PLUS, 0, 0} +} + +func (d *ServerCoreData) Serialize() []byte { + return []byte{} +} + +func (d *ServerCoreData) ScType() Message { + return SC_CORE +} +func (d *ServerCoreData) Unpack(r io.Reader) error { + version, _ := core.ReadUInt32LE(r) + d.RdpVersion = VERSION(version) + d.ClientRequestedProtocol, _ = core.ReadUInt32LE(r) + d.EarlyCapabilityFlags, _ = core.ReadUInt32LE(r) + + return nil + //return struc.Unpack(r, d) +} + +type ServerNetworkData struct { + MCSChannelId uint16 `struc:"little"` + ChannelCount uint16 `struc:"little,sizeof=ChannelIdArray"` + ChannelIdArray []uint16 `struc:"little"` +} + +func NewServerNetworkData() *ServerNetworkData { + return &ServerNetworkData{} +} +func (d *ServerNetworkData) ScType() Message { + return SC_NET +} +func (d *ServerNetworkData) Unpack(r io.Reader) error { + return struc.Unpack(r, d) +} + +type CertData interface { + GetPublicKey() (*rsa.PublicKey, error) + Verify() bool + Unpack(io.Reader) error +} +type ServerCertificate struct { + DwVersion uint32 + CertData CertData +} + +func (sc *ServerCertificate) Unpack(r io.Reader) error { + sc.DwVersion, _ = core.ReadUInt32LE(r) + var cd CertData + switch CertificateType(sc.DwVersion & 0x7fffffff) { + case CERT_CHAIN_VERSION_1: + slog.Debug("ProprietaryServerCertificate") + cd = &ProprietaryServerCertificate{} + case CERT_CHAIN_VERSION_2: + slog.Debug("X509CertificateChain") + cd = &X509CertificateChain{} + default: + slog.Error("Unpack", "Unsupported version", sc.DwVersion&0x7fffffff) + return errors.New("Unsupported version") + } + if cd != nil { + err := cd.Unpack(r) + if err != nil { + slog.Error("Unpack", "err", err) + return err + } + } + sc.CertData = cd + + return nil +} + +type ServerSecurityData struct { + EncryptionMethod uint32 `struc:"little"` + EncryptionLevel uint32 `struc:"little"` + ServerRandomLen uint32 //0x00000020 + ServerCertLen uint32 + ServerRandom []byte + ServerCertificate ServerCertificate +} + +func NewServerSecurityData() *ServerSecurityData { + return &ServerSecurityData{ + 0, 0, 0x00000020, 0, []byte{}, ServerCertificate{}} +} +func (d *ServerSecurityData) ScType() Message { + return SC_SECURITY +} +func (s *ServerSecurityData) Unpack(r io.Reader) error { + s.EncryptionMethod, _ = core.ReadUInt32LE(r) + s.EncryptionLevel, _ = core.ReadUInt32LE(r) + if !(s.EncryptionMethod == 0 && s.EncryptionLevel == 0) { + s.ServerRandomLen, _ = core.ReadUInt32LE(r) + s.ServerCertLen, _ = core.ReadUInt32LE(r) + s.ServerRandom, _ = core.ReadBytes(int(s.ServerRandomLen), r) + var sc ServerCertificate + data, _ := core.ReadBytes(int(s.ServerCertLen), r) + rd := bytes.NewReader(data) + err := sc.Unpack(rd) + if err != nil { + return err + } + s.ServerCertificate = sc + } + + return nil +} + +// ServerMsgChannelData holds the message channel ID allocated by the server +// (TS_UD_SC_MCS_MSGCHANNEL). Used for connect-time network auto-detection. +type ServerMsgChannelData struct { + MCSChannelId uint16 +} + +func (d *ServerMsgChannelData) ScType() Message { + return SC_MCS_MSGCHANNEL +} + +func (d *ServerMsgChannelData) Unpack(r io.Reader) error { + var err error + d.MCSChannelId, err = core.ReadUint16LE(r) + return err +} + +// PackClientMsgChannelData serialises the TS_UD_CS_MCS_MSGCHANNEL block. +// Advertising this block requests that the server allocate a dedicated message +// channel for connect-time network auto-detection (RTT/BW measurements). +func PackClientMsgChannelData() []byte { + buff := &bytes.Buffer{} + core.WriteUInt16LE(CS_MCS_MSGCHANNEL, buff) // type 0xC006 + core.WriteUInt16LE(0x08, buff) // length = 8 + core.WriteUInt32LE(0, buff) // flags = 0 + return buff.Bytes() +} + +func MakeConferenceCreateRequest(userData []byte) []byte { + buff := &bytes.Buffer{} + per.WriteChoice(0, buff) // 00 + per.WriteObjectIdentifier(t124_02_98_oid, buff) // 05:00:14:7c:00:01 + per.WriteLength(len(userData)+14, buff) + per.WriteChoice(0, buff) // 00 + per.WriteSelection(0x08, buff) // 08 + per.WriteNumericString("1", 1, buff) // 00 10 + per.WritePadding(1, buff) // 00 + per.WriteNumberOfSet(1, buff) // 01 + per.WriteChoice(0xc0, buff) // c0 + per.WriteOctetStream(h221_cs_key, 4, buff) // 00 44:75:63:61 + per.WriteOctetStream(string(userData), 0, buff) + return buff.Bytes() +} + +type ScData interface { + ScType() Message + Unpack(io.Reader) error +} + +func ReadConferenceCreateResponse(data []byte) []any { + ret := make([]any, 0, 3) + + r := bytes.NewReader(data) + per.ReadChoice(r) + if !per.ReadObjectIdentifier(r, t124_02_98_oid) { + slog.Error("NODE_RDP_PROTOCOL_T125_GCC_BAD_OBJECT_IDENTIFIER_T124") + return ret + } + per.ReadLength(r) + per.ReadChoice(r) + per.ReadInteger16(r) + per.ReadInteger(r) + per.ReadEnumerates(r) + per.ReadNumberOfSet(r) + per.ReadChoice(r) + + if !per.ReadOctetStream(r, h221_sc_key, 4) { + slog.Error("NODE_RDP_PROTOCOL_T125_GCC_BAD_H221_SC_KEY") + return ret + } + + ln, _ := per.ReadLength(r) + for ln > 0 { + t, _ := core.ReadUint16LE(r) + l, _ := core.ReadUint16LE(r) + dataBytes, _ := core.ReadBytes(int(l)-4, r) + ln = ln - l + var d ScData + switch Message(t) { + case SC_CORE: + d = &ServerCoreData{} + case SC_SECURITY: + d = &ServerSecurityData{} + case SC_NET: + d = &ServerNetworkData{} + case SC_MCS_MSGCHANNEL: + d = &ServerMsgChannelData{} + default: + slog.Debug("ReadConferenceCreateResponse: ignoring unknown block", "type", t) + continue + } + + if d != nil { + r := bytes.NewReader(dataBytes) + err := d.Unpack(r) + if err != nil { + slog.Warn("ReadConferenceCreateResponse", "err", err) + } + ret = append(ret, d) + } + } + + return ret +} diff --git a/protocol/t125/mcs.go b/protocol/t125/mcs.go new file mode 100644 index 0000000..f75d82c --- /dev/null +++ b/protocol/t125/mcs.go @@ -0,0 +1,753 @@ +package t125 + +import ( + "bytes" + "errors" + "fmt" + "io" + "log/slog" + "reflect" + "time" + + // "git.zeroonesoft.cn/golib/rdplib/plugin/cliprdr" + "git.zeroonesoft.cn/golib/rdplib/plugin/drdynvc" + "git.zeroonesoft.cn/golib/rdplib/plugin/rail" + "git.zeroonesoft.cn/golib/rdplib/plugin/rdpsnd" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/ber" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/gcc" + "git.zeroonesoft.cn/golib/rdplib/protocol/t125/per" +) + +// take idea from https://github.com/Madnikulin50/gordp + +// Multiple Channel Service layer + +type MCSMessage uint8 + +const ( + MCS_TYPE_CONNECT_INITIAL MCSMessage = 0x65 + MCS_TYPE_CONNECT_RESPONSE = 0x66 +) + +type MCSDomainPDU uint16 + +const ( + ERECT_DOMAIN_REQUEST MCSDomainPDU = 1 + DISCONNECT_PROVIDER_ULTIMATUM = 8 + ATTACH_USER_REQUEST = 10 + ATTACH_USER_CONFIRM = 11 + CHANNEL_JOIN_REQUEST = 14 + CHANNEL_JOIN_CONFIRM = 15 + SEND_DATA_REQUEST = 25 + SEND_DATA_INDICATION = 26 +) + +const ( + MCS_GLOBAL_CHANNEL_ID uint16 = 1003 + MCS_USERCHANNEL_BASE = 1001 +) + +const ( + GLOBAL_CHANNEL_NAME = "global" +) + +/** + * Format MCS PDULayer header packet + * @param mcsPdu {integer} + * @param options {integer} + * @returns {type.UInt8} headers + */ +func writeMCSPDUHeader(mcsPdu MCSDomainPDU, options uint8, w io.Writer) { + core.WriteUInt8((uint8(mcsPdu)<<2)|options, w) +} + +func readMCSPDUHeader(options uint8, mcsPdu MCSDomainPDU) bool { + return (options >> 2) == uint8(mcsPdu) +} + +type DomainParameters struct { + MaxChannelIds int + MaxUserIds int + MaxTokenIds int + NumPriorities int + MinThoughput int + MaxHeight int + MaxMCSPDUsize int + ProtocolVersion int +} + +/** + * @see http://www.itu.int/rec/T-REC-T.125-199802-I/en page 25 + * @returns {asn1.univ.Sequence} + */ +func NewDomainParameters( + maxChannelIds int, + maxUserIds int, + maxTokenIds int, + numPriorities int, + minThoughput int, + maxHeight int, + maxMCSPDUsize int, + protocolVersion int) *DomainParameters { + return &DomainParameters{maxChannelIds, maxUserIds, maxTokenIds, + numPriorities, minThoughput, maxHeight, maxMCSPDUsize, protocolVersion} +} + +func (d *DomainParameters) BER() []byte { + buff := &bytes.Buffer{} + ber.WriteInteger(d.MaxChannelIds, buff) + ber.WriteInteger(d.MaxUserIds, buff) + ber.WriteInteger(d.MaxTokenIds, buff) + ber.WriteInteger(1, buff) + ber.WriteInteger(0, buff) + ber.WriteInteger(1, buff) + ber.WriteInteger(d.MaxMCSPDUsize, buff) + ber.WriteInteger(2, buff) + return buff.Bytes() +} + +func ReadDomainParameters(r io.Reader) (*DomainParameters, error) { + if !ber.ReadUniversalTag(ber.TAG_SEQUENCE, true, r) { + return nil, errors.New("bad BER tags") + } + d := &DomainParameters{} + ber.ReadLength(r) + + d.MaxChannelIds, _ = ber.ReadInteger(r) + d.MaxUserIds, _ = ber.ReadInteger(r) + d.MaxTokenIds, _ = ber.ReadInteger(r) + ber.ReadInteger(r) + ber.ReadInteger(r) + ber.ReadInteger(r) + d.MaxMCSPDUsize, _ = ber.ReadInteger(r) + ber.ReadInteger(r) + return d, nil +} + +/** + * @see http://www.itu.int/rec/T-REC-T.125-199802-I/en page 25 + * @param userData {Buffer} + * @returns {asn1.univ.Sequence} + */ +type ConnectInitial struct { + CallingDomainSelector []byte + CalledDomainSelector []byte + UpwardFlag bool + TargetParameters DomainParameters + MinimumParameters DomainParameters + MaximumParameters DomainParameters + UserData []byte +} + +func NewConnectInitial(userData []byte) ConnectInitial { + return ConnectInitial{[]byte{0x1}, + []byte{0x1}, + true, + *NewDomainParameters(34, 2, 0, 1, 0, 1, 0xffff, 2), + *NewDomainParameters(1, 1, 1, 1, 0, 1, 0x420, 2), + *NewDomainParameters(0xffff, 0xfc17, 0xffff, 1, 0, 1, 0xffff, 2), + userData} +} + +func (c *ConnectInitial) BER() []byte { + buff := &bytes.Buffer{} + ber.WriteOctetstring(string(c.CallingDomainSelector), buff) + ber.WriteOctetstring(string(c.CalledDomainSelector), buff) + ber.WriteBoolean(c.UpwardFlag, buff) + ber.WriteEncodedDomainParams(c.TargetParameters.BER(), buff) + ber.WriteEncodedDomainParams(c.MinimumParameters.BER(), buff) + ber.WriteEncodedDomainParams(c.MaximumParameters.BER(), buff) + ber.WriteOctetstring(string(c.UserData), buff) + return buff.Bytes() +} + +/** + * @see http://www.itu.int/rec/T-REC-T.125-199802-I/en page 25 + * @returns {asn1.univ.Sequence} + */ + +type ConnectResponse struct { + result uint8 + calledConnectId int + domainParameters *DomainParameters + userData []byte +} + +func NewConnectResponse(userData []byte) *ConnectResponse { + return &ConnectResponse{0, + 0, + NewDomainParameters(22, 3, 0, 1, 0, 1, 0xfff8, 2), + userData} +} + +func ReadConnectResponse(r io.Reader) (*ConnectResponse, error) { + c := &ConnectResponse{} + var err error + _, err = ber.ReadApplicationTag(MCS_TYPE_CONNECT_RESPONSE, r) + if err != nil { + return nil, err + } + c.result, err = ber.ReadEnumerated(r) + if err != nil { + return nil, err + } + + c.calledConnectId, err = ber.ReadInteger(r) + c.domainParameters, err = ReadDomainParameters(r) + if err != nil { + return nil, err + } + if !ber.ReadUniversalTag(ber.TAG_OCTET_STRING, false, r) { + return nil, errors.New("invalid expected BER tag") + } + dataLen, _ := ber.ReadLength(r) + c.userData, err = core.ReadBytes(dataLen, r) + return c, err +} + +type MCSChannelInfo struct { + ID uint16 + Name string +} + +type MCS struct { + emission.Emitter + transport core.Transport + recvOpCode MCSDomainPDU + sendOpCode MCSDomainPDU + channels []MCSChannelInfo +} + +func NewMCS(t core.Transport, recvOpCode MCSDomainPDU, sendOpCode MCSDomainPDU) *MCS { + m := &MCS{ + *emission.NewEmitter(), + t, + recvOpCode, + sendOpCode, + []MCSChannelInfo{{MCS_GLOBAL_CHANNEL_ID, GLOBAL_CHANNEL_NAME}}, + } + + m.transport.On("close", func() { + m.Emit("close") + }).On("error", func(err error) { + m.Emit("error", err) + }) + return m +} + +func (x *MCS) Read(b []byte) (n int, err error) { + return x.transport.Read(b) +} + +func (x *MCS) Write(b []byte) (n int, err error) { + return x.transport.Write(b) +} + +func (m *MCS) Close() error { + return m.transport.Close() +} + +type MCSClient struct { + *MCS + clientCoreData *gcc.ClientCoreData + clientNetworkData *gcc.ClientNetworkData + clientSecurityData *gcc.ClientSecurityData + // lowColorDepth:SetSessionColorDepth(16/24) 置位, Dynvc-GFX 能力位 + // 需要跳过(该端点要求 32bpp 会话) + lowColorDepth bool + + serverCoreData *gcc.ServerCoreData + serverNetworkData *gcc.ServerNetworkData + serverSecurityData *gcc.ServerSecurityData + + channelsConnected int + userId uint16 + nbChannelRequested int + pendingJoins int // 并行突发加入后尚未收到确认的 SVC 通道数 + messageChannelId uint16 // from SC_MCS_MSGCHANNEL; 0 = not negotiated + messageChannelJoined bool + bwStartTime time.Time // timestamp of last RDP_BW_START for timeDelta calculation +} + +func NewMCSClient(t core.Transport, kbdLayout uint32, keyboardType uint32, keyboardSubType uint32) *MCSClient { + c := &MCSClient{ + MCS: NewMCS(t, SEND_DATA_INDICATION, SEND_DATA_REQUEST), + clientCoreData: gcc.NewClientCoreData(kbdLayout, keyboardType, keyboardSubType), + clientNetworkData: gcc.NewClientNetworkData(), + clientSecurityData: gcc.NewClientSecurityData(), + userId: 1 + MCS_USERCHANNEL_BASE, + } + c.transport.On("connect", c.connect) + return c +} + +func (c *MCSClient) SetClientDesktop(width, height uint16) { + c.clientCoreData.DesktopWidth = width + c.clientCoreData.DesktopHeight = height +} + +// SetClientName 覆盖客户端计算机名(ClientCoreData.ClientName,UTF-16LE, +// 字段共 32 字节,超长截断)。服务端按该名字管理 \\tsclient 重定向映射, +// 同名客户端残留(非正常断开)会让后续同名会话的映射失效——Explorer +// 打开 \\tsclient\<名> 报“试图访问无效的地址”且 rdpdr 通道零 IRP +// (RDPDR-2)。wasm 下 os.Hostname 回退值恒为 "js",必须每次连接随机化。 +func (c *MCSClient) SetClientName(name string) { + var buf [32]byte + copy(buf[:], core.UnicodeEncode(name)) + c.clientCoreData.ClientName = buf +} + +// SetSessionColorDepth 请求会话颜色位数(16/24/32,其它值按 32 处理)。 +// 32bpp = highColorDepth 24BPP + WANT_32BPP_SESSION(mstsc 默认); +// 16/24bpp 清除 WANT_32BPP_SESSION 并收紧 supportedColorDepths。 +// 注意 RDPGFX 会话(RemoteFX/H264 模式)表面恒为 32bpp,此选项只在 +// 传统位图管线生效。 +// SetSessionColorDepth 请求会话颜色位数(16/24/32,其它值按 32 处理)。 +// 实测(Win10 19041)三个坑: +// 1. supportedColorDepths 若不含 24/32bpp 支持位,服务器在 MCS 握手 +// 阶段直接断连(无 ERRINFO)——支持位恒为全量(与 FreeRDP 一致)。 +// 2. 低色深时必须同时清除 SUPPORT_DYNVC_GFX_PROTOCOL(0x200): +// 该位声明 32bpp 的 Dynvc-GFX 端点,与 16bpp 会话矛盾,服务器会 +// 进入不一致状态(fastpath SURFCMDS 泛滥、音频 DVC 不初始化)。 +// SetClientDynvcProtocol 会在低色深时跳过该位(drdynvc 通道保留)。 +// 3. 色深协商只通过 highColorDepth + postBeta2ColorDepth + +// WANT_32BPP_SESSION 位表达。 +// 注意 RDPGFX 会话(RemoteFX/H264 模式)表面恒为 32bpp,此选项只在 +// 传统位图管线生效。 +func (c *MCSClient) SetSessionColorDepth(bpp int) { + d := c.clientCoreData + switch bpp { + case 16: + d.HighColorDepth = gcc.HIGH_COLOR_16BPP + d.PostBeta2ColorDepth = gcc.RNS_UD_COLOR_16BPP_565 + d.EarlyCapabilityFlags &^= gcc.RNS_UD_CS_WANT_32BPP_SESSION | + gcc.RNS_UD_CS_SUPPORT_DYNVC_GFX_PROTOCOL + c.lowColorDepth = true + case 24: + d.HighColorDepth = gcc.HIGH_COLOR_24BPP + d.PostBeta2ColorDepth = gcc.RNS_UD_COLOR_24BPP + d.EarlyCapabilityFlags &^= gcc.RNS_UD_CS_WANT_32BPP_SESSION | + gcc.RNS_UD_CS_SUPPORT_DYNVC_GFX_PROTOCOL + c.lowColorDepth = true + default: + // 32bpp:保持 NewClientCoreData 的默认(WANT_32BPP_SESSION) + } +} + +func (c *MCSClient) SetClientDynvcProtocol() { + // 低色深会话不支持 Dynvc-GFX 端点(见 SetSessionColorDepth), + // drdynvc 通道本身保留(剪贴板图片/音频 DVC 仍需要)。 + if !c.lowColorDepth { + c.clientCoreData.EarlyCapabilityFlags |= gcc.RNS_UD_CS_SUPPORT_DYNVC_GFX_PROTOCOL + } + c.clientNetworkData.AddVirtualChannel(drdynvc.ChannelName, drdynvc.ChannelOption) +} + +func (c *MCSClient) SetClientRemoteProgram() { + c.clientNetworkData.AddVirtualChannel(rail.ChannelName, rail.ChannelOption) +} + +func (c *MCSClient) SetClientSoundProtocol() { + c.clientNetworkData.AddVirtualChannel(rdpsnd.ChannelName, rdpsnd.ChannelOption) +} + +func (c *MCSClient) SetClientDeviceRedirection() { + c.clientNetworkData.AddVirtualChannel("rdpdr", + uint32(gcc.CHANNEL_OPTION_INITIALIZED|gcc.CHANNEL_OPTION_ENCRYPT_RDP|gcc.CHANNEL_OPTION_COMPRESS_RDP)) +} + +func (c *MCSClient) SetClientClipboard() { + c.clientNetworkData.AddVirtualChannel("cliprdr", + uint32(gcc.CHANNEL_OPTION_INITIALIZED|gcc.CHANNEL_OPTION_ENCRYPT_RDP|gcc.CHANNEL_OPTION_COMPRESS_RDP)) +} + +func (c *MCSClient) connect(selectedProtocol uint32) { + slog.Debug("connect", "selectedProtocol", selectedProtocol) + c.clientCoreData.ServerSelectedProtocol = selectedProtocol + + slog.Debug("connnect", "clientCoreData", c.clientCoreData) + slog.Debug("connect", "clientNetworkData", c.clientNetworkData) + slog.Debug("connect", "clientSecurityData", c.clientSecurityData) + // sendConnectclientCoreDataInitial + userDataBuff := bytes.Buffer{} + userDataBuff.Write(c.clientCoreData.Pack()) + userDataBuff.Write(c.clientNetworkData.Pack()) + userDataBuff.Write(c.clientSecurityData.Pack()) + userDataBuff.Write(gcc.PackClientMsgChannelData()) + + slog.Debug("userData", "data", core.Hex(userDataBuff.Bytes()), "len", len(userDataBuff.Bytes())) + ccReq := gcc.MakeConferenceCreateRequest(userDataBuff.Bytes()) + slog.Debug("ccReq", "data", core.Hex(ccReq), "len", len(ccReq)) + connectInitial := NewConnectInitial(ccReq) + connectInitialBerEncoded := connectInitial.BER() + + dataBuff := &bytes.Buffer{} + ber.WriteApplicationTag(uint8(MCS_TYPE_CONNECT_INITIAL), len(connectInitialBerEncoded), dataBuff) + dataBuff.Write(connectInitialBerEncoded) + slog.Debug("send connet initial", "data", core.Hex(dataBuff.Bytes()), "len", len(dataBuff.Bytes())) + + _, err := c.transport.Write(dataBuff.Bytes()) + if err != nil { + c.Emit("error", errors.New(fmt.Sprintf("mcs sendConnectInitial write error %v", err))) + return + } + slog.Debug("mcs wait for data event") + c.transport.Once("data", c.recvConnectResponse) +} + +func (c *MCSClient) recvConnectResponse(s []byte) { + slog.Debug("mcs recvConnectResponse", "s", core.Hex(s)) + cResp, err := ReadConnectResponse(bytes.NewReader(s)) + if err != nil { + c.Emit("error", errors.New(fmt.Sprintf("ReadConnectResponse %v", err))) + return + } + // record server gcc block + serverSettings := gcc.ReadConferenceCreateResponse(cResp.userData) + for _, v := range serverSettings { + switch v.(type) { + case *gcc.ServerSecurityData: + c.serverSecurityData = v.(*gcc.ServerSecurityData) + + case *gcc.ServerCoreData: + c.serverCoreData = v.(*gcc.ServerCoreData) + + case *gcc.ServerNetworkData: + c.serverNetworkData = v.(*gcc.ServerNetworkData) + + case *gcc.ServerMsgChannelData: + c.messageChannelId = v.(*gcc.ServerMsgChannelData).MCSChannelId + slog.Debug("SC_MCS_MSGCHANNEL", "messageChannelId", c.messageChannelId) + + default: + slog.Warn("recvConnectResponse: unhandled server gcc block", "type", reflect.TypeOf(v)) + } + } + c.sendErectDomainRequest() + c.sendAttachUserRequest() + + c.transport.Once("data", c.recvAttachUserConfirm) +} + +func (c *MCSClient) sendErectDomainRequest() { + buff := &bytes.Buffer{} + writeMCSPDUHeader(ERECT_DOMAIN_REQUEST, 0, buff) + per.WriteInteger(0, buff) + per.WriteInteger(0, buff) + c.transport.Write(buff.Bytes()) +} + +func (c *MCSClient) sendAttachUserRequest() { + buff := &bytes.Buffer{} + writeMCSPDUHeader(ATTACH_USER_REQUEST, 0, buff) + c.transport.Write(buff.Bytes()) +} + +func (c *MCSClient) recvAttachUserConfirm(s []byte) { + slog.Debug("mcs recvAttachUserConfirm", "s", core.Hex(s)) + r := bytes.NewReader(s) + + option, err := core.ReadUInt8(r) + if err != nil { + c.Emit("error", err) + return + } + + if !readMCSPDUHeader(option, ATTACH_USER_CONFIRM) { + c.Emit("error", errors.New("NODE_RDP_PROTOCOL_T125_MCS_BAD_HEADER")) + return + } + + e, err := per.ReadEnumerates(r) + if err != nil { + c.Emit("error", err) + return + } + if e != 0 { + c.Emit("error", errors.New("NODE_RDP_PROTOCOL_T125_MCS_SERVER_REJECT_USER'")) + return + } + + userId, _ := per.ReadInteger16(r) + userId += MCS_USERCHANNEL_BASE + c.userId = userId + + c.channels = append(c.channels, MCSChannelInfo{userId, "user"}) + c.connectChannels() +} + +func (c *MCSClient) connectChannels() { + slog.Debug("connectChannels", "channelsConnected", c.channelsConnected, "channels", c.channels) + if c.channelsConnected < len(c.channels) { + // sendChannelJoinRequest + c.sendChannelJoinRequest(c.channels[c.channelsConnected].ID) + + c.transport.Once("data", c.recvChannelJoinConfirm) + return + } + + // Join the message channel (for connect-time auto-detection) before SVCs. + if c.messageChannelId != 0 && !c.messageChannelJoined { + c.messageChannelJoined = true + c.sendChannelJoinRequest(c.messageChannelId) + c.transport.Once("data", c.recvChannelJoinConfirm) + return + } + + if c.nbChannelRequested == 0 && int(c.serverNetworkData.ChannelCount) > 0 { + // 并行突发加入全部 SVC 静态通道(mstsc 同款):原先串行逐个等 + // 确认,rdpdr 排在队尾,其 join 在服务端 LogonNotify 之后 ~100ms + // 才完成——登录期驱动映射(winlogon→drprov 按"当时的设备表"创建 + // \\TSCLIENT\<名> 连接)因此永远看不到我们的设备,`dir + // \\tsclient\<名>` 恒为"设备没有连接"(2250)。并行发送让 rdpdr + // 与其余通道同时就绪,抢在 LogonNotify 之前。 + for i := 0; i < int(c.serverNetworkData.ChannelCount); i++ { + c.sendChannelJoinRequest(c.serverNetworkData.ChannelIdArray[i]) + } + c.nbChannelRequested = int(c.serverNetworkData.ChannelCount) + c.pendingJoins = int(c.serverNetworkData.ChannelCount) + // 单个持久监听 + 计数:emission 对同一条数据事件会触发全部 + // 监听器,N×Once 会在首个确认上重复消费(实测导致通道表重复 + // 追加与 MCS opcode 错误),故用 Off 在计数归零后摘除。 + c.transport.On("data", c.recvChannelJoinConfirm) + return + } + if c.pendingJoins > 0 { + // 并行突发确认进行中,由 recvChannelJoinConfirm 计数收尾。 + return + } + c.transport.On("data", c.recvData) + // send client and sever gcc informations callback to sec + clientData := make([]any, 0) + clientData = append(clientData, c.clientCoreData) + clientData = append(clientData, c.clientSecurityData) + clientData = append(clientData, c.clientNetworkData) + + serverData := make([]any, 0) + serverData = append(serverData, c.serverCoreData) + serverData = append(serverData, c.serverSecurityData) + c.Emit("connect", clientData, serverData, c.userId, c.channels) +} + +func (c *MCSClient) sendChannelJoinRequest(channelId uint16) { + slog.Debug("sendChannelJoinRequest", "channelId", channelId) + buff := &bytes.Buffer{} + writeMCSPDUHeader(CHANNEL_JOIN_REQUEST, 0, buff) + per.WriteInteger16(c.userId-MCS_USERCHANNEL_BASE, buff) + per.WriteInteger16(channelId, buff) + c.transport.Write(buff.Bytes()) +} + +func (c *MCSClient) recvData(s []byte) { + r := bytes.NewReader(s) + option, err := core.ReadUInt8(r) + if err != nil { + c.Emit("error", err) + return + } + + if readMCSPDUHeader(option, DISCONNECT_PROVIDER_ULTIMATUM) { + c.Emit("error", errors.New("MCS DISCONNECT_PROVIDER_ULTIMATUM")) + c.transport.Close() + return + } else if !readMCSPDUHeader(option, c.recvOpCode) { + c.Emit("error", errors.New("Invalid expected MCS opcode receive data")) + return + } + + userId, _ := per.ReadInteger16(r) + userId += MCS_USERCHANNEL_BASE + + channelId, _ := per.ReadInteger16(r) + per.ReadEnumerates(r) + size, _ := per.ReadLength(r) + // channel ID doesn't match a requested layer + found := false + channelName := "" + for _, channel := range c.channels { + if channel.ID == channelId { + found = true + channelName = channel.Name + break + } + } + if !found { + if c.messageChannelId != 0 && channelId == c.messageChannelId { + data, _ := core.ReadBytes(int(size), r) + c.handleAutoDetect(data) + return + } + slog.Error("mcs receive data for an unconnected layer") + return + } + left, err := core.ReadBytes(int(size), r) + if err != nil { + c.Emit("error", errors.New(fmt.Sprintf("mcs recvData get data error %v", err))) + return + } + c.Emit("sec", channelName, left) +} + +func (c *MCSClient) recvChannelJoinConfirm(s []byte) { + slog.Debug("recvChannelJoinConfirm", "s", core.Hex(s)) + r := bytes.NewReader(s) + option, err := core.ReadUInt8(r) + if err != nil { + return + } + + if !readMCSPDUHeader(option, CHANNEL_JOIN_CONFIRM) { + // 并行突发窗口内同监听器会看到数据 PDU:静默忽略(由 recvData + // 在突发完成后接管处理)。 + return + } + + confirm, _ := per.ReadEnumerates(r) + userId, _ := per.ReadInteger16(r) + userId += MCS_USERCHANNEL_BASE + + if c.userId != userId { + c.Emit("error", errors.New("NODE_RDP_PROTOCOL_T125_MCS_INVALID_USER_ID")) + return + } + + channelId, _ := per.ReadInteger16(r) + if (confirm != 0) && (channelId == uint16(MCS_GLOBAL_CHANNEL_ID) || channelId == c.userId) { + c.Emit("error", errors.New("NODE_RDP_PROTOCOL_T125_MCS_SERVER_MUST_CONFIRM_STATIC_CHANNEL")) + return + } + if confirm == 0 { + for i := 0; i < int(c.serverNetworkData.ChannelCount); i++ { + if channelId == c.serverNetworkData.ChannelIdArray[i] { + var t MCSChannelInfo + t.ID = channelId + t.Name = string(c.clientNetworkData.ChannelDefArray[i].Name[:]) + c.channels = append(c.channels, t) + } + } + } + c.channelsConnected++ + if c.pendingJoins > 0 { + c.pendingJoins-- + if c.pendingJoins > 0 { + return + } + // 全部 SVC 确认到齐:摘除突发期监听器,交还数据泵。 + c.transport.Off("data", c.recvChannelJoinConfirm) + } + c.connectChannels() +} + +// Connect-time auto-detection constants (MS-RDPBCGR 2.2.14). +const ( + secAutoDetectReq = uint16(0x1000) + secAutoDetectRsp = uint16(0x2000) + + rdpRttRequestConnecttime = uint16(0x1001) + rdpBwStartConnecttime = uint16(0x1014) + rdpBwPayload = uint16(0x0002) + rdpBwStopConnecttime = uint16(0x002B) + rdpRttRequest = uint16(0x0001) // continuous RTT request + rdpBwStart = uint16(0x0014) // continuous BW start (no response) + rdpBwStop = uint16(0x0429) // continuous BW stop + + typeIDAutodetectResponse = uint8(0x01) + rdpRttResponseType = uint16(0x0000) + rdpBwResultsConnecttime = uint16(0x0003) + rdpBwResults = uint16(0x000B) // continuous BW results +) + +// handleAutoDetect processes a connect-time auto-detect request from the server +// on the message channel. It responds to RTT and BW measurement requests so +// that gnome-remote-desktop proceeds to open the audio DVC channels. +func (c *MCSClient) handleAutoDetect(data []byte) { + r := bytes.NewReader(data) + secFlag, _ := core.ReadUint16LE(r) + core.ReadUint16LE(r) // secFlagHi + + if secFlag&secAutoDetectReq == 0 { + return + } + + _, _ = core.ReadUInt8(r) // headerLength + core.ReadUInt8(r) // headerTypeId + seqNum, _ := core.ReadUint16LE(r) + reqType, _ := core.ReadUint16LE(r) + + switch reqType { + case rdpRttRequestConnecttime, rdpRttRequest: + c.sendAutoDetectResponse(seqNum, rdpRttResponseType, 0) + case rdpBwStartConnecttime, rdpBwStart: + c.bwStartTime = time.Now() + case rdpBwStopConnecttime: + elapsed := uint32(time.Since(c.bwStartTime).Milliseconds()) + if elapsed == 0 { + elapsed = 1 + } + c.sendAutoDetectResponse(seqNum, rdpBwResultsConnecttime, elapsed) + case rdpBwStop: + elapsed := uint32(time.Since(c.bwStartTime).Milliseconds()) + if elapsed == 0 { + elapsed = 1 + } + c.sendAutoDetectResponse(seqNum, rdpBwResults, elapsed) + // rdpBwPayload requires no response + } +} + +// sendAutoDetectResponse sends an auto-detect response on the message channel. +// timeDelta is 0 for RTT responses (no BW fields); non-zero for BW responses +// (timeDelta in milliseconds since the corresponding BW_START was received). +func (c *MCSClient) sendAutoDetectResponse(sequenceNumber uint16, responseType uint16, timeDelta uint32) { + includeBW := responseType == rdpBwResultsConnecttime || responseType == rdpBwResults + headerLength := uint8(6) + if includeBW { + headerLength = 14 + } + + payload := &bytes.Buffer{} + core.WriteUInt16LE(secAutoDetectRsp, payload) + core.WriteUInt16LE(0, payload) + core.WriteUInt8(headerLength, payload) + core.WriteUInt8(typeIDAutodetectResponse, payload) + core.WriteUInt16LE(sequenceNumber, payload) + core.WriteUInt16LE(responseType, payload) + if includeBW { + core.WriteUInt32LE(timeDelta, payload) // timeDelta in milliseconds + core.WriteUInt32LE(0, payload) // byteCount (no BW_PAYLOAD was sent) + } + + c.transport.Write(c.Pack(payload.Bytes(), c.messageChannelId)) +} + +func (c *MCSClient) Pack(data []byte, channelId uint16) []byte { + buff := &bytes.Buffer{} + writeMCSPDUHeader(c.sendOpCode, 0, buff) + per.WriteInteger16(c.userId-MCS_USERCHANNEL_BASE, buff) + per.WriteInteger16(channelId, buff) + core.WriteUInt8(0x70, buff) + per.WriteLength(len(data), buff) + core.WriteBytes(data, buff) + return buff.Bytes() +} + +func (c *MCSClient) Write(data []byte) (n int, err error) { + data = c.Pack(data, c.channels[0].ID) + return c.transport.Write(data) +} + +func (c *MCSClient) SendToChannel(channel string, data []byte) (n int, err error) { + channelId := c.channels[0].ID + for _, ch := range c.channels { + if channel == ch.Name { + channelId = ch.ID + break + } + } + + data = c.Pack(data, channelId) + return c.transport.Write(data) +} diff --git a/protocol/t125/per/per.go b/protocol/t125/per/per.go new file mode 100644 index 0000000..8861de9 --- /dev/null +++ b/protocol/t125/per/per.go @@ -0,0 +1,211 @@ +package per + +import ( + "io" + "log/slog" + + "git.zeroonesoft.cn/golib/rdplib/core" +) + +// zeroPad is a shared zero buffer used by WritePadding to avoid per-call heap +// allocations. Writing is done in chunks up to len(zeroPad). +var zeroPad [256]byte + +func ReadEnumerates(r io.Reader) (uint8, error) { + return core.ReadUInt8(r) +} + +func WriteInteger(n int, w io.Writer) { + if n <= 0xff { + WriteLength(1, w) + core.WriteUInt8(uint8(n), w) + } else if n <= 0xffff { + WriteLength(2, w) + core.WriteUInt16BE(uint16(n), w) + } else { + WriteLength(4, w) + core.WriteUInt32BE(uint32(n), w) + } +} + +func ReadInteger16(r io.Reader) (uint16, error) { + return core.ReadUint16BE(r) +} + +func WriteInteger16(value uint16, w io.Writer) { + core.WriteUInt16BE(value, w) +} + +/** + * @param choice {integer} + * @returns {type.UInt8} choice per encoded + */ +func WriteChoice(choice uint8, w io.Writer) { + core.WriteUInt8(choice, w) +} + +/** + * @param value {raw} value to convert to per format + * @returns type objects per encoding value + */ +func WriteLength(value int, w io.Writer) { + if value > 0x7f { + core.WriteUInt16BE(uint16(value|0x8000), w) + } else { + core.WriteUInt8(uint8(value), w) + } +} + +func ReadLength(r io.Reader) (uint16, error) { + b, err := core.ReadUInt8(r) + if err != nil { + return 0, nil + } + var size uint16 + if b&0x80 > 0 { + b = b &^ 0x80 + size = uint16(b) << 8 + left, _ := core.ReadUInt8(r) + size += uint16(left) + } else { + size = uint16(b) + } + return size, nil +} + +/** + * @param oid {array} oid to write + * @returns {type.Component} per encoded object identifier + */ +func WriteObjectIdentifier(oid []byte, w io.Writer) { + core.WriteUInt8(5, w) + core.WriteByte((oid[0]<<4)&(oid[1]&0x0f), w) + core.WriteByte(oid[2], w) + core.WriteByte(oid[3], w) + core.WriteByte(oid[4], w) + core.WriteByte(oid[5], w) +} + +/** + * @param selection {integer} + * @returns {type.UInt8} per encoded selection + */ +func WriteSelection(selection uint8, w io.Writer) { + core.WriteUInt8(selection, w) +} + +func WriteNumericString(s string, minValue int, w io.Writer) { + length := len(s) + mLength := minValue + if length >= minValue { + mLength = length - minValue + } + WriteLength(mLength, w) + var buf [1]byte + for i := 0; i < length; i += 2 { + c1 := int(s[i]) + c2 := 0x30 + if i+1 < length { + c2 = int(s[i+1]) + } + c1 = (c1 - 0x30) % 10 + c2 = (c2 - 0x30) % 10 + buf[0] = uint8((c1 << 4) | c2) + w.Write(buf[:]) + } +} + +func WritePadding(length int, w io.Writer) { + for length > 0 { + n := length + if n > len(zeroPad) { + n = len(zeroPad) + } + w.Write(zeroPad[:n]) + length -= n + } +} + +func WriteNumberOfSet(n int, w io.Writer) { + core.WriteUInt8(uint8(n), w) +} + +/** + * @param oStr {String} + * @param minValue {integer} default 0 + * @returns {type.Component} per encoded octet stream + */ +func WriteOctetStream(oStr string, minValue int, w io.Writer) { + length := len(oStr) + mlength := minValue + + if length-minValue >= 0 { + mlength = length - minValue + } + WriteLength(mlength, w) + io.WriteString(w, oStr) +} + +func ReadChoice(r io.Reader) uint8 { + choice, _ := core.ReadUInt8(r) + return choice +} +func ReadNumberOfSet(r io.Reader) uint8 { + choice, _ := core.ReadUInt8(r) + return choice +} +func ReadInteger(r io.Reader) uint32 { + size, _ := ReadLength(r) + switch size { + case 1: + ret, _ := core.ReadUInt8(r) + return uint32(ret) + case 2: + ret, _ := core.ReadUint16BE(r) + return uint32(ret) + case 4: + ret, _ := core.ReadUInt32BE(r) + return ret + default: + slog.Debug("ReadInteger") + } + return 0 +} + +func ReadObjectIdentifier(r io.Reader, oid []byte) bool { + size, _ := ReadLength(r) + if size != 5 { + return false + } + + a_oid := []byte{0, 0, 0, 0, 0, 0} + t12, _ := core.ReadByte(r) + a_oid[0] = t12 >> 4 + a_oid[1] = t12 & 0x0f + a_oid[2], _ = core.ReadByte(r) + a_oid[3], _ = core.ReadByte(r) + a_oid[4], _ = core.ReadByte(r) + a_oid[5], _ = core.ReadByte(r) + + for i := range oid { + if oid[i] != a_oid[i] { + return false + } + } + return true +} +func ReadOctetStream(r io.Reader, s string, min int) bool { + ln, _ := ReadLength(r) + size := int(ln) + min + if size != len(s) { + return false + } + for i := range size { + b, _ := core.ReadByte(r) + if b != s[i] { + return false + } + } + + return true +} diff --git a/protocol/tpkt/tpkt.go b/protocol/tpkt/tpkt.go new file mode 100644 index 0000000..7bf32e8 --- /dev/null +++ b/protocol/tpkt/tpkt.go @@ -0,0 +1,293 @@ +package tpkt + +import ( + "encoding/binary" + "fmt" + "io" + "log/slog" + "runtime/debug" + "strings" + "sync" + "time" + + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" + "git.zeroonesoft.cn/golib/rdplib/protocol/nla" +) + +var writePool = sync.Pool{ + New: func() any { return make([]byte, 0, 4096) }, +} + +// readBufPool reuses packet body buffers in readLoop. +// Typical RDP packets are well under 4 KiB; the pool avoids a heap +// allocation for every incoming packet (~60/s during active sessions). +// Buffers larger than maxPooledReadBuf are not pooled to avoid keeping +// large slices alive in the pool between bursts. +const maxPooledReadBuf = 32 * 1024 + +var readBufPool = sync.Pool{ + New: func() any { return make([]byte, 0, 4096) }, +} + +func acquireReadBuf(n int) []byte { + b := readBufPool.Get().([]byte) + if cap(b) >= n { + return b[:n] + } + readBufPool.Put(b[:0]) + return make([]byte, n) +} + +func releaseReadBuf(b []byte) { + if cap(b) <= maxPooledReadBuf { + readBufPool.Put(b[:0]) + } +} + +// take idea from https://github.com/Madnikulin50/gordp + +/** + * Type of tpkt packet + * Fastpath is use to shortcut RDP stack + * @see http://msdn.microsoft.com/en-us/library/cc240621.aspx + * @see http://msdn.microsoft.com/en-us/library/cc240589.aspx + */ +const ( + FASTPATH_ACTION_FASTPATH = 0x0 + FASTPATH_ACTION_X224 = 0x3 +) + +/** + * TPKT layer of rdp stack + */ +type TPKT struct { + emission.Emitter + Conn *core.SocketLayer + ntlm *nla.NTLMv2 + fastPathListener core.FastPathListener + ntlmSec *nla.NTLMv2Security +} + +func New(s *core.SocketLayer, ntlm *nla.NTLMv2) *TPKT { + t := &TPKT{ + Emitter: *emission.NewEmitter(), + Conn: s, + ntlm: ntlm, + } + go t.readLoop() + return t +} + +// readLoop is the single goroutine that reads all incoming TPKT/FastPath packets. +// It replaces the previous callback-chain pattern (StartReadBytes → recvHeader → +// StartReadBytes → recvExtendedHeader → …) which spawned a new goroutine for each +// individual read. By using a single blocking loop with io.ReadFull, we eliminate +// goroutine creation/destruction overhead on the hot receive path. +func (t *TPKT) readLoop() { + var hdr [2]byte + for { + if _, err := io.ReadFull(t.Conn, hdr[:]); err != nil { + t.Emit("error", err) + return + } + + version := hdr[0] + if version == FASTPATH_ACTION_X224 { + // TPKT packet: 4-byte header total (version, reserved, length-hi, length-lo) + var extHdr [2]byte + if _, err := io.ReadFull(t.Conn, extHdr[:]); err != nil { + t.Emit("error", err) + return + } + size := binary.BigEndian.Uint16(extHdr[:]) + if size < 4 { + t.Emit("error", fmt.Errorf("TPKT: invalid packet size %d", size)) + return + } + body := acquireReadBuf(int(size) - 4) + if _, err := io.ReadFull(t.Conn, body); err != nil { + t.Emit("error", err) + return + } + t.Emit("data", body) + releaseReadBuf(body) + } else { + // FastPath packet: 2- or 3-byte header + secFlag := (version >> 6) & 0x3 + length := int(hdr[1]) + slog.Debug("TPKT FastPath", "secFlag", secFlag, "length", length) + + var packetSize int + if length&0x80 != 0 { + // Extended 3-byte header: high 7 bits from hdr[1], low 8 from next byte + var extByte [1]byte + if _, err := io.ReadFull(t.Conn, extByte[:]); err != nil { + slog.Error("TPKT recvExtendedFastPathHeader", "err", err) + return + } + leftPart := length & ^0x80 + packetSize = (leftPart<<8) + int(extByte[0]) - 3 + } else { + packetSize = length - 2 + } + + if packetSize < 0 { + t.Emit("error", fmt.Errorf("TPKT FastPath: invalid packet size %d", packetSize)) + return + } + body := acquireReadBuf(packetSize) + if _, err := io.ReadFull(t.Conn, body); err != nil { + slog.Debug("TPKT recvFastPath error", "err", err) + return + } + t.recvFastPathSafe(secFlag, body) + releaseReadBuf(body) + } + } +} + +func (t *TPKT) StartTLS() error { + return t.Conn.StartTLS() +} + +// recvFastPathSafe 隔离 fast-path 解码 panic:解码器遇到意外数据时记录 +// 堆栈并上报错误横幅,而不是让整个 Go 运行时崩溃(worker 只会显示 +// "Go program has already exited",毫无诊断信息)。 +func (t *TPKT) recvFastPathSafe(secFlag byte, body []byte) { + defer func() { + if r := recover(); r != nil { + st := strings.Split(string(debug.Stack()), "\n") + if len(st) > 14 { + st = st[:14] + } + slog.Error("RecvFastPath panic", "recover", r, "stack", strings.Join(st, "\n")) + t.Emit("error", fmt.Errorf("fast-path panic: %v\n%s", r, strings.Join(st, "\n"))) + } + }() + t.fastPathListener.RecvFastPath(secFlag, body) +} + +func (t *TPKT) StartNLA() error { + // Set a deadline for the entire NLA handshake (TLS + NTLM auth) + // to prevent hanging indefinitely when the server is slow to respond. + t.Conn.SetDeadline(time.Now().Add(30 * time.Second)) + defer t.Conn.SetDeadline(time.Time{}) // clear deadline after NLA completes + + slog.Debug("StartNLA: TLS handshake begin") + err := t.StartTLS() + if err != nil { + slog.Error("StartNLA", "start tls failed", err) + return err + } + slog.Debug("StartNLA: TLS handshake complete") + req := nla.EncodeDERTRequest([]nla.Message{t.ntlm.GetNegotiateMessage()}, nil, nil) + slog.Debug("StartNLA send", "req", core.Hex(req), "len", len(req)) + _, err = t.Conn.Write(req) + if err != nil { + slog.Error("send NegotiateMessage", "err", err) + return err + } + + resp := make([]byte, 1024) + n, err := t.Conn.Read(resp) + slog.Debug("StartNLA recv", "n", n, "err", err) + if err != nil { + return fmt.Errorf("read %s", err) + } else { + slog.Debug("StartNLA Read success") + } + return t.recvChallenge(resp[:n]) +} + +func (t *TPKT) recvChallenge(data []byte) error { + slog.Debug("recvChallenge", "data", core.Hex(data)) + tsreq, err := nla.DecodeDERTRequest(data) + if err != nil { + slog.Debug("DecodeDERTRequest", "err", err) + return err + } + slog.Debug("recvChallenge", "tsreq", tsreq) + // get pubkey + pubkey, err := t.Conn.TlsPubKey() + slog.Debug("recvChallenge", "pubkey", core.Hex(pubkey)) + + authMsg, ntlmSec := t.ntlm.GetAuthenticateMessage(tsreq.NegoTokens[0].Data) + t.ntlmSec = ntlmSec + + encryptPubkey := ntlmSec.GssEncrypt(pubkey) + req := nla.EncodeDERTRequest([]nla.Message{authMsg}, nil, encryptPubkey) + slog.Debug("recvChallenge", "send", core.Hex(req), "len", len(req)) + _, err = t.Conn.Write(req) + if err != nil { + slog.Error("send AuthenticateMessage", "err", err) + return err + } + + slog.Debug("recvChallenge read challenge start") + resp := make([]byte, 1024) + n, err := t.Conn.Read(resp) + if err != nil { + slog.Error("recvChallenge", "err", err) + return fmt.Errorf("read %s", err) + } + + return t.recvPubKeyInc(resp[:n]) +} + +func (t *TPKT) recvPubKeyInc(data []byte) error { + slog.Debug("recvPubKeyInc", "data", core.Hex(data), "len", len(data)) + tsreq, err := nla.DecodeDERTRequest(data) + if err != nil { + slog.Debug("DecodeDERTRequest", "err", err) + return err + } + slog.Debug("PubKeyAuth", "key", core.Hex(tsreq.PubKeyAuth)) + //ignore + pubkey := t.ntlmSec.GssDecrypt([]byte(tsreq.PubKeyAuth)) + slog.Debug("GssDecrypy", "pubkey", core.Hex(pubkey)) + domain, username, password := t.ntlm.GetEncodedCredentials() + credentials := nla.EncodeDERTCredentials(domain, username, password) + authInfo := t.ntlmSec.GssEncrypt(credentials) + req := nla.EncodeDERTRequest(nil, authInfo, nil) + _, err = t.Conn.Write(req) + if err != nil { + slog.Debug("send AuthenticateMessage", "err", err) + return err + } + + return nil +} + +func (t *TPKT) Read(b []byte) (n int, err error) { + return t.Conn.Read(b) +} + +func (t *TPKT) Write(data []byte) (n int, err error) { + buf := writePool.Get().([]byte) + size := uint16(len(data) + 4) + buf = append(buf[:0], FASTPATH_ACTION_X224, 0, byte(size>>8), byte(size)) + buf = append(buf, data...) + n, err = t.Conn.Write(buf) + writePool.Put(buf[:0]) + return +} + +func (t *TPKT) Close() error { + return t.Conn.Close() +} + +func (t *TPKT) SetFastPathListener(f core.FastPathListener) { + t.fastPathListener = f +} + +func (t *TPKT) SendFastPath(secFlag byte, data []byte) (n int, err error) { + buf := writePool.Get().([]byte) + hdr := uint16(len(data)+3) | 0x8000 + buf = append(buf[:0], FASTPATH_ACTION_FASTPATH|((secFlag&0x3)<<6), byte(hdr>>8), byte(hdr)) + buf = append(buf, data...) + n, err = t.Conn.Write(buf) + writePool.Put(buf[:0]) + return +} + diff --git a/protocol/x224/x224.go b/protocol/x224/x224.go new file mode 100644 index 0000000..3985c16 --- /dev/null +++ b/protocol/x224/x224.go @@ -0,0 +1,351 @@ +package x224 + +import ( + "bytes" + "errors" + "fmt" + "log/slog" + "sync" + + "github.com/lunixbochs/struc" + "git.zeroonesoft.cn/golib/rdplib/core" + "git.zeroonesoft.cn/golib/rdplib/emission" + "git.zeroonesoft.cn/golib/rdplib/protocol/tpkt" +) + +var x224WritePool = sync.Pool{ + New: func() any { return &bytes.Buffer{} }, +} + +// take idea from https://github.com/Madnikulin50/gordp + +/** + * Message type present in X224 packet header + */ +type MessageType byte + +const ( + TPDU_CONNECTION_REQUEST MessageType = 0xE0 + TPDU_CONNECTION_CONFIRM = 0xD0 + TPDU_DISCONNECT_REQUEST = 0x80 + TPDU_DATA = 0xF0 + TPDU_ERROR = 0x70 +) + +/** + * Type of negotiation present in negotiation packet + */ +type NegotiationType byte + +const ( + TYPE_RDP_NEG_REQ NegotiationType = 0x01 + TYPE_RDP_NEG_RSP = 0x02 + TYPE_RDP_NEG_FAILURE = 0x03 +) + +/** + * Protocols available for x224 layer + */ + +const ( + PROTOCOL_RDP uint32 = 0x00000000 + PROTOCOL_SSL = 0x00000001 + PROTOCOL_HYBRID = 0x00000002 + PROTOCOL_HYBRID_EX = 0x00000008 +) + +/** + * Use to negotiate security layer of RDP stack + * In node-rdpjs only ssl is available + * @param opt {object} component type options + * @see request -> http://msdn.microsoft.com/en-us/library/cc240500.aspx + * @see response -> http://msdn.microsoft.com/en-us/library/cc240506.aspx + * @see failure ->http://msdn.microsoft.com/en-us/library/cc240507.aspx + */ +type Negotiation struct { + Type NegotiationType `struc:"byte"` + Flag uint8 `struc:"uint8"` + Length uint16 `struc:"little"` + Result uint32 `struc:"little"` +} + +func NewNegotiation() *Negotiation { + return &Negotiation{0, 0, 0x0008 /*constant*/, PROTOCOL_RDP} +} + +type failureCode int + +const ( + //The server requires that the client support Enhanced RDP Security (section 5.4) with either TLS 1.0, 1.1 or 1.2 (section 5.4.5.1) or CredSSP (section 5.4.5.2). If only CredSSP was requested then the server only supports TLS. + SSL_REQUIRED_BY_SERVER = 0x00000001 + + //The server is configured to only use Standard RDP Security mechanisms (section 5.3) and does not support any External Security Protocols (section 5.4.5). + SSL_NOT_ALLOWED_BY_SERVER = 0x00000002 + + //The server does not possess a valid authentication certificate and cannot initialize the External Security Protocol Provider (section 5.4.5). + SSL_CERT_NOT_ON_SERVER = 0x00000003 + + //The list of requested security protocols is not consistent with the current security protocol in effect. This error is only possible when the Direct Approach (sections 5.4.2.2 and 1.3.1.2) is used and an External Security Protocol (section 5.4.5) is already being used. + INCONSISTENT_FLAGS = 0x00000004 + + //The server requires that the client support Enhanced RDP Security (section 5.4) with CredSSP (section 5.4.5.2). + HYBRID_REQUIRED_BY_SERVER = 0x00000005 + + //The server requires that the client support Enhanced RDP Security (section 5.4) with TLS 1.0, 1.1 or 1.2 (section 5.4.5.1) and certificate-based client authentication.<4> + SSL_WITH_USER_AUTH_REQUIRED_BY_SERVER = 0x00000006 +) + +/** + * X224 client connection request + * @param opt {object} component type options + * @see http://msdn.microsoft.com/en-us/library/cc240470.aspx + */ +type ClientConnectionRequestPDU struct { + Len uint8 + Code MessageType + Padding1 uint16 + Padding2 uint16 + Padding3 uint8 + Cookie []byte + requestedProtocol uint32 + ProtocolNeg *Negotiation +} + +func NewClientConnectionRequestPDU(cookie []byte, requestedProtocol uint32) *ClientConnectionRequestPDU { + x := ClientConnectionRequestPDU{0, TPDU_CONNECTION_REQUEST, 0, 0, 0, + cookie, requestedProtocol, NewNegotiation()} + + x.Len = 6 + if len(cookie) > 0 { + x.Len += uint8(len(cookie) + 2) + } + if x.requestedProtocol > PROTOCOL_RDP { + x.Len += 8 + } + + return &x +} + +func (x *ClientConnectionRequestPDU) Serialize() []byte { + buff := &bytes.Buffer{} + core.WriteUInt8(x.Len, buff) + core.WriteUInt8(uint8(x.Code), buff) + core.WriteUInt16BE(x.Padding1, buff) + core.WriteUInt16BE(x.Padding2, buff) + core.WriteUInt8(x.Padding3, buff) + + if len(x.Cookie) > 0 { + buff.Write(x.Cookie) + core.WriteUInt8(0x0D, buff) + core.WriteUInt8(0x0A, buff) + } + + if x.requestedProtocol > PROTOCOL_RDP { + struc.Pack(buff, x.ProtocolNeg) + } + + return buff.Bytes() +} + +/** + * X224 Server connection confirm + * @param opt {object} component type options + * @see http://msdn.microsoft.com/en-us/library/cc240506.aspx + */ +type ServerConnectionConfirm struct { + Len uint8 + Code MessageType + Padding1 uint16 + Padding2 uint16 + Padding3 uint8 + ProtocolNeg *Negotiation +} + +/** + * Header of each data message from x224 layer + * @returns {type.Component} + */ +type DataHeader struct { + Header uint8 `struc:"little"` + MessageType MessageType `struc:"uint8"` + Separator uint8 `struc:"little"` +} + +func NewDataHeader() *DataHeader { + return &DataHeader{2, TPDU_DATA /* constant */, 0x80 /*constant*/} +} + +/** + * Common X224 Automata + * @param presentation {Layer} presentation layer + */ +type X224 struct { + emission.Emitter + transport core.Transport + requestedProtocol uint32 + selectedProtocol uint32 + dataHeader *DataHeader + username string + routingToken []byte +} + +func New(t core.Transport) *X224 { + x := &X224{ + Emitter: *emission.NewEmitter(), + transport: t, + requestedProtocol: PROTOCOL_RDP | PROTOCOL_SSL | PROTOCOL_HYBRID, + selectedProtocol: PROTOCOL_SSL, + dataHeader: NewDataHeader(), + } + + t.On("close", func() { + x.Emit("close") + }).On("error", func(err error) { + x.Emit("error", err) + }) + + return x +} + +func (x *X224) Read(b []byte) (n int, err error) { + return x.transport.Read(b) +} + +func (x *X224) Write(b []byte) (n int, err error) { + buff := x224WritePool.Get().(*bytes.Buffer) + buff.Reset() + err = struc.Pack(buff, x.dataHeader) + if err != nil { + x224WritePool.Put(buff) + return 0, err + } + buff.Write(b) + n, err = x.transport.Write(buff.Bytes()) + x224WritePool.Put(buff) + return +} + +func (x *X224) Close() error { + return x.transport.Close() +} + +func (x *X224) SetRequestedProtocol(p uint32) { + x.requestedProtocol = p +} + +func (x *X224) SetUsername(username string) { + x.username = username +} + +func (x *X224) SetRoutingToken(token []byte) { + x.routingToken = token +} + +func (x *X224) Connect() error { + if x.transport == nil { + return errors.New("no transport") + } + + var cookie string + if len(x.routingToken) > 0 { + // Use server-provided routing token (strip trailing \r\n if present) + tok := x.routingToken + if len(tok) >= 2 && tok[len(tok)-2] == 0x0D && tok[len(tok)-1] == 0x0A { + tok = tok[:len(tok)-2] + } + cookie = string(tok) + } else { + name := x.username + if name == "" { + name = "test" + } + cookie = "Cookie: mstshash=" + name + } + + message := NewClientConnectionRequestPDU([]byte(cookie), x.requestedProtocol) + message.ProtocolNeg.Type = TYPE_RDP_NEG_REQ + message.ProtocolNeg.Result = uint32(x.requestedProtocol) + + slog.Debug("x224 Connect", "message", core.Hex(message.Serialize())) + _, err := x.transport.Write(message.Serialize()) + x.transport.Once("data", x.recvConnectionConfirm) + return err +} + +func (x *X224) recvConnectionConfirm(s []byte) { + slog.Debug("x224 recvConnectionConfirm", "s", core.Hex(s)) + r := bytes.NewReader(s) + ln, _ := core.ReadUInt8(r) + if ln > 6 { + message := &ServerConnectionConfirm{} + if err := struc.Unpack(bytes.NewReader(s), message); err != nil { + slog.Error("ReadServerConnectionConfirm", "err", err) + x.Emit("error", err) + return + } + slog.Debug("recvConnectionConfirm", "message", *message.ProtocolNeg) + if message.ProtocolNeg.Type == TYPE_RDP_NEG_FAILURE { + negErr := fmt.Errorf("NODE_RDP_PROTOCOL_X224_NEG_FAILURE with code: %d, see https://msdn.microsoft.com/en-us/library/cc240507.aspx", + message.ProtocolNeg.Result) + slog.Error(negErr.Error()) + //only use Standard RDP Security mechanisms + if message.ProtocolNeg.Result == 2 { + slog.Debug("Only use Standard RDP Security mechanisms, Reconnect with Standard RDP") + } + x.Emit("error", negErr) + x.Close() + return + } + + if message.ProtocolNeg.Type == TYPE_RDP_NEG_RSP { + slog.Debug("TYPE_RDP_NEG_RSP") + x.selectedProtocol = message.ProtocolNeg.Result + } + } else { + x.selectedProtocol = PROTOCOL_RDP + } + + if x.selectedProtocol == PROTOCOL_HYBRID_EX { + err := errors.New("NODE_RDP_PROTOCOL_HYBRID_EX_NOT_SUPPORTED") + slog.Error(err.Error()) + x.Emit("error", err) + return + } + + x.transport.On("data", x.recvData) + + if x.selectedProtocol == PROTOCOL_RDP { + slog.Debug("*** RDP security selected ***") + x.Emit("connect", x.selectedProtocol) + return + } + + if x.selectedProtocol == PROTOCOL_SSL { + slog.Debug("*** SSL security selected ***") + err := x.transport.(*tpkt.TPKT).StartTLS() + if err != nil { + slog.Error("start tls failed:", "err", err) + x.Emit("error", err) + return + } + x.Emit("connect", x.selectedProtocol) + return + } + + if x.selectedProtocol == PROTOCOL_HYBRID { + slog.Debug("*** NLA Security selected ***") + err := x.transport.(*tpkt.TPKT).StartNLA() + if err != nil { + slog.Error("start NLA failed:", "err", err) + x.Emit("error", err) + return + } + x.Emit("connect", x.selectedProtocol) + return + } +} + +func (x *X224) recvData(s []byte) { + // x224 header takes 3 bytes + x.Emit("data", s[3:]) +}